mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 08:36:57 +02:00
Compare commits
37 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 7fe4fab167 | |||
| 37b924945d | |||
| 5631346b07 | |||
| 97eb3c635a | |||
| 05f7923860 | |||
| 49028cb755 | |||
| adf71d1458 | |||
| cefda6524c | |||
| d0f14de739 | |||
| 9e7646ee62 | |||
| 6e6cfc89db | |||
| 3719f135e3 | |||
| 2724b4a96c | |||
| e386e290ab | |||
| 0a44376403 | |||
| 9e61269935 | |||
| a0836aa819 | |||
| cf2800f7bd | |||
| fc89b9d14c | |||
| 3e5326e60a | |||
| 840841a53c | |||
| cfb4ab24c3 | |||
| 67f9cfad91 | |||
| 1c78f5e500 | |||
| ad8ff2e45e | |||
| c13a1ff4ec | |||
| 8b66ebaf72 | |||
| bbdaf5c3f3 | |||
| a515aeaf39 | |||
| 0f6f14eaf6 | |||
| dafdd34af9 | |||
| f6791130df | |||
| 145a6267fa | |||
| 7716f1da23 | |||
| beb1d7d89f | |||
| fa0d593f28 | |||
| ff48040c78 |
@@ -12,9 +12,9 @@ jobs:
|
||||
steps:
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
- uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version: '1.26'
|
||||
go-version: '1.25'
|
||||
check-latest: true
|
||||
|
||||
- name: Build
|
||||
@@ -38,9 +38,9 @@ jobs:
|
||||
steps:
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
- uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version: '1.26'
|
||||
go-version: '1.25'
|
||||
check-latest: true
|
||||
|
||||
- name: Build
|
||||
@@ -78,9 +78,9 @@ jobs:
|
||||
steps:
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
- uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version: '1.26'
|
||||
go-version: '1.25'
|
||||
check-latest: true
|
||||
|
||||
- name: Import certificates
|
||||
|
||||
@@ -32,9 +32,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
- uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version: '1.26'
|
||||
go-version: '1.25'
|
||||
check-latest: true
|
||||
|
||||
- name: add hashicorp source
|
||||
@@ -64,9 +64,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
- uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version: '1.26'
|
||||
go-version: '1.25'
|
||||
check-latest: true
|
||||
|
||||
- name: add hashicorp source
|
||||
@@ -90,9 +90,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
- uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version: '1.26'
|
||||
go-version: '1.25'
|
||||
check-latest: true
|
||||
|
||||
# WSL2 + Ubuntu so the smoke can run a real linux peer with its own
|
||||
|
||||
@@ -20,9 +20,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
- uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version: '1.26'
|
||||
go-version: '1.25'
|
||||
check-latest: true
|
||||
|
||||
- name: build
|
||||
|
||||
@@ -20,9 +20,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
- uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version: '1.26'
|
||||
go-version: '1.25'
|
||||
check-latest: true
|
||||
|
||||
- name: Install goimports
|
||||
@@ -42,7 +42,7 @@ jobs:
|
||||
- name: golangci-lint
|
||||
uses: golangci/golangci-lint-action@v9
|
||||
with:
|
||||
version: v2.12
|
||||
version: v2.5
|
||||
|
||||
test:
|
||||
name: Test ${{ matrix.name }}
|
||||
@@ -80,9 +80,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
- uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version: '1.26'
|
||||
go-version: '1.25'
|
||||
check-latest: true
|
||||
|
||||
- name: Build
|
||||
@@ -125,9 +125,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
- uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version: '1.26'
|
||||
go-version: '1.25'
|
||||
check-latest: true
|
||||
|
||||
- name: Build ${{ matrix.name }}
|
||||
|
||||
@@ -7,88 +7,6 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
||||
|
||||
## [Unreleased]
|
||||
|
||||
## [1.11.0] - 2026-07-23
|
||||
|
||||
See the [v1.11.0](https://github.com/slackhq/nebula/milestone/25?closed=1) milestone for a complete list of changes.
|
||||
|
||||
### Breaking
|
||||
|
||||
- Logging has switched from logrus to Go's structured `slog`. Log output changes: levels are upper case
|
||||
(`level=INFO`), trace prints as `level=DEBUG-4`, timestamps are always RFC3339Nano and `logging.timestamp_format`
|
||||
is ignored, and some messages were reworded. Review any log parsing before upgrading. This is also an API break
|
||||
for embedders, as constructors now take a `*slog.Logger`. (#1672, #1734, #1621)
|
||||
- `firewall.inbound_action` and `firewall.outbound_action` (used to set reject vs. drop policy) were each being
|
||||
applied to the opposite direction, that is now corrected. This only affects how blocked packets are answered, not
|
||||
which packets the firewall allows or denies. If you set either of these you are getting the behavior of the other
|
||||
one today and likely want to swap them before upgrading. (#1798)
|
||||
- On Windows, Nebula now installs WFP PERMIT filters for the nebula adapter and the listener port by default. WFP
|
||||
sits below Windows Defender Firewall, so any WDF inbound rules you rely on for either will no longer apply. Set
|
||||
`tun.windows_bypass_wdf` and `listen.windows_bypass_wdf` to false to leave WDF in charge. (#1710)
|
||||
- On Windows, the nebula device is now set to the `private` network category instead of whatever Windows decided,
|
||||
which is usually `Public`. This makes the host firewall less restrictive on the overlay. Set
|
||||
`tun.network_category` to `unset` to keep the old behavior. (#1710)
|
||||
- Reject packets for non-TCP now use ICMP code 13, communication administratively prohibited, instead of code 3,
|
||||
port unreachable. Anything keying off the old code needs updating. (#1766, #1768)
|
||||
- The SSH debug server's profiling commands are now confined to `sshd.sandbox_dir`, which defaults to
|
||||
`$TMP/nebula-debug`. Relative paths resolve inside it and absolute paths outside it are rejected, so anything
|
||||
scripting `start-cpu-profile`, `save-heap-profile`, or `save-mutex-profile` with a path elsewhere needs the
|
||||
directory set. The directory is not created for you. (#1622)
|
||||
|
||||
### Added
|
||||
|
||||
- Sign the Windows release binaries. (#1718)
|
||||
- Generate IPv6 reject packets, matching the existing IPv4 behavior. (#1766, #1767, #1768)
|
||||
- Accept `-` in `nebula-cert` to read from stdin or write to stdout. (#1714)
|
||||
- Search for both `config.yml` and `config.yaml` in service and command line modes. (#1717)
|
||||
- Add version labels to the Docker/OCI images. (#1772)
|
||||
- Rebind the listener and re-query lighthouses on macOS when the underlay network changes, so devices moving
|
||||
between wifi and wired or between networks recover without waiting for dead tunnel detection. Controlled by
|
||||
`listen.rebind_on_network_change` (default `true`, not reloadable). (#1816)
|
||||
|
||||
### Changed
|
||||
|
||||
- Reload the firewall when the unsafe networks in the certificate change. (#1719)
|
||||
- Reconfigure, start, and stop the stats listener on a config reload instead of requiring a restart. (#1670)
|
||||
- Update a static host's addresses when they change on reload. (#1713)
|
||||
- Don't require a port on ICMP firewall rules. (#1609)
|
||||
- Connection track ICMP traffic. (#1602)
|
||||
- Return `NODATA` instead of `NXDOMAIN` from the DNS server for a name that exists but has no record of the
|
||||
requested type, so clients that query `AAAA` first (busybox/Alpine) fall through to `A`. (#1668)
|
||||
- Record the local host's details in the DNS server. (#1716)
|
||||
- Install Windows unsafe routes as link routes. (#1709)
|
||||
- Reduce relay handshake log spam, and only log a handshake send error at error level when the remote list
|
||||
changes. (#1733, #1765, #1810)
|
||||
- Start, stop, and reload subsystems (DNS, stats, conntrack, ssh, punchy) cleanly without leaking goroutines. (#1640, #1654, #1661, #1667, #1669, #1708, #1806, #1815)
|
||||
- `Control` is now safe to stop and wait on from any lifecycle state, and a new `Control.Wait` blocks until nebula
|
||||
has fully stopped and returns the first fatal reader error. Failed starts release the udp sockets and tun fd
|
||||
instead of leaking them. (#1794)
|
||||
- Trigger an immediate lighthouse update when reconnecting to or adding a lighthouse instead of waiting for the next update tick. (#1645)
|
||||
- Bring the Darwin and OpenBSD tun implementations in line with the other BSDs. (#1703)
|
||||
- Update to build against go v1.26. (#1818)
|
||||
- Various dependency updates. (#1586, #1587, #1604, #1617, #1618, #1627, #1628, #1629, #1652, #1664, #1665, #1697, #1721, #1732, #1742, #1743, #1750, #1763, #1771, #1782, #1800, #1807)
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix a data race on a host's remote address that could send packets to the wrong address during a roam. (#1773)
|
||||
- Fix tunnels that could permanently escape connection manager monitoring. (#1752)
|
||||
- Fix a crash when reloading the SSH server's trusted keys. (#1787)
|
||||
- Fix hostmap corruption when a host has multiple overlay addresses. Each address now gets its own list instead of
|
||||
a single shared chain, which also fixes two latent bugs on the add and makePrimary paths. (#1788, #1790)
|
||||
- Apply `remote_allow_list` IPv4 rules to 4-in-6 mapped addresses. (#1786)
|
||||
- Don't panic in the DNS server on a short or empty query name. (#1635)
|
||||
- Advance the replay window on relayed packets so a relay drops replayed frames instead of re-forwarding them. (#1751)
|
||||
- Fix a race in relay state handling. (#1753)
|
||||
- Lock replay window updates so concurrent readers can't corrupt it. (#1802)
|
||||
- Reject malformed handshakes more reliably, including invalid ed25519 key lengths. (#1601, #1756)
|
||||
- Properly handle `closetunnel` packets. (#1638)
|
||||
- Fix an IPv6 extension-header length overflow that could make the firewall parse the wrong protocol and ports. (#1789)
|
||||
- Fix relay re-establishment when a handshake arrives over a relay entry that a one-sided teardown left
|
||||
`Disestablished`, which silently dropped every send until dead tunnel detection forced a re-handshake. (#1805)
|
||||
- Don't build new relay state on a tunnel that was just discarded. (#1796)
|
||||
- Don't delete the wrong pending hostinfo in the handshake manager. (#1811)
|
||||
- Don't call the packet reader after a UDP error on Darwin. (#1755)
|
||||
- Open the FreeBSD tun device non blocking. (#1666)
|
||||
|
||||
## [1.10.3] - 2026-02-06
|
||||
|
||||
### Security
|
||||
|
||||
@@ -1,96 +0,0 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
|
||||
package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
cert_test "github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestControlStopClosesOnTimer reproduces the dnclient lifecycle: nebula runs as
|
||||
// a library, and on a config update dnclient calls Stop() in-process to tear the
|
||||
// old instance down before starting a new one. This boots a real nebula (real
|
||||
// blocking UDP sockets, tun disabled), lets it run, then Stop()s it on a timer
|
||||
// and asserts it actually closes. If the reader goroutines parked in recvmmsg
|
||||
// don't wake on Close(), Wait() blocks forever and this fails with a goroutine
|
||||
// dump instead of relying on a process signal to unstick them.
|
||||
func TestControlStopClosesOnTimer(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
dir := t.TempDir()
|
||||
|
||||
before := time.Now().Add(-time.Hour)
|
||||
after := time.Now().Add(time.Hour)
|
||||
ca, _, caKey, caPEM := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, before, after, nil, nil, nil)
|
||||
networks := []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")}
|
||||
_, _, keyPEM, certPEM := cert_test.NewTestCert(cert.Version2, cert.Curve_CURVE25519, ca, caKey, "close-on-timer", before, after, networks, nil, nil)
|
||||
|
||||
caPath := filepath.Join(dir, "ca.pem")
|
||||
certPath := filepath.Join(dir, "cert.pem")
|
||||
keyPath := filepath.Join(dir, "key.pem")
|
||||
require.NoError(t, os.WriteFile(caPath, caPEM, 0o600))
|
||||
require.NoError(t, os.WriteFile(certPath, certPEM, 0o600))
|
||||
require.NoError(t, os.WriteFile(keyPath, keyPEM, 0o600))
|
||||
|
||||
// tun disabled so no device/root is needed; routines: 2 so we exercise the
|
||||
// multi-socket (SO_REUSEPORT) teardown, which is where dnclient runs.
|
||||
configBody := fmt.Sprintf(`
|
||||
pki:
|
||||
ca: %s
|
||||
cert: %s
|
||||
key: %s
|
||||
listen:
|
||||
host: 127.0.0.1
|
||||
port: 0
|
||||
tun:
|
||||
disabled: true
|
||||
firewall:
|
||||
outbound:
|
||||
- port: any
|
||||
proto: any
|
||||
host: any
|
||||
inbound:
|
||||
- port: any
|
||||
proto: any
|
||||
host: any
|
||||
routines: 2
|
||||
`, caPath, certPath, keyPath)
|
||||
require.NoError(t, os.WriteFile(filepath.Join(dir, "config.yml"), []byte(configBody), 0o600))
|
||||
|
||||
c := config.NewC(l)
|
||||
require.NoError(t, c.Load(dir))
|
||||
|
||||
ctrl, err := nebula.Main(c, false, "close-on-timer", l, nil)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, ctrl.Start())
|
||||
|
||||
// Run like a live nebula, then close on a timer, exactly as dnclient does.
|
||||
<-time.NewTimer(5 * time.Second).C
|
||||
|
||||
stopped := make(chan struct{})
|
||||
go func() {
|
||||
ctrl.Stop() // closes the udp sockets (shutdown(2)) and the tun
|
||||
ctrl.Wait() // blocks until every reader goroutine has returned
|
||||
close(stopped)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-stopped:
|
||||
t.Log("nebula closed cleanly on timer")
|
||||
case <-time.After(10 * time.Second):
|
||||
buf := make([]byte, 1<<20)
|
||||
n := runtime.Stack(buf, true)
|
||||
t.Fatalf("nebula did NOT close within 10s of Stop(): a blocking reader never woke\n%s", buf[:n])
|
||||
}
|
||||
}
|
||||
@@ -25,7 +25,6 @@ func newTestLighthouse() *LightHouse {
|
||||
lighthouses := []netip.Addr{}
|
||||
staticList := map[netip.Addr]struct{}{}
|
||||
|
||||
lh.localAddrsFn = func(*LocalAllowList) []netip.Addr { return nil }
|
||||
lh.lighthouses.Store(&lighthouses)
|
||||
lh.staticList.Store(&staticList)
|
||||
|
||||
|
||||
@@ -2,26 +2,16 @@ package nebula
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/handshake"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/noiseutil"
|
||||
)
|
||||
|
||||
const ReplayWindow = 8192
|
||||
|
||||
// sessionEpoch hands out a receiver-local ordinal to every ConnectionState at creation. The RX
|
||||
// staging sort (overlay/batch) orders packets by (epoch, message counter). A re-handshake never
|
||||
// rekeys an existing tunnel; it brings up a new hostinfo and ConnectionState with a counter space
|
||||
// starting near zero, while the old tunnel keeps decrypting until torn down. During that cutover
|
||||
// one flush batch can hold packets from both tunnels, and the epoch keeps the old tunnel's
|
||||
// packets sorted first.
|
||||
var sessionEpoch atomic.Uint64
|
||||
|
||||
type ConnectionState struct {
|
||||
eKey noiseutil.CipherState
|
||||
dKey noiseutil.CipherState
|
||||
@@ -30,10 +20,7 @@ type ConnectionState struct {
|
||||
initiator bool
|
||||
messageCounter atomic.Uint64
|
||||
window *Bits
|
||||
decryptLock sync.Mutex
|
||||
writeLock sync.Mutex
|
||||
// epoch is this session's sessionEpoch ordinal. Immutable after creation.
|
||||
epoch uint64
|
||||
}
|
||||
|
||||
// newConnectionStateFromResult builds a fully-populated ConnectionState from a
|
||||
@@ -48,7 +35,6 @@ func newConnectionStateFromResult(r *handshake.Result) *ConnectionState {
|
||||
eKey: noiseutil.NewCipherState(r.EKey, r.Cipher),
|
||||
dKey: noiseutil.NewCipherState(r.DKey, r.Cipher),
|
||||
window: NewBits(ReplayWindow),
|
||||
epoch: sessionEpoch.Add(1),
|
||||
}
|
||||
ci.messageCounter.Add(r.MessageIndex)
|
||||
for i := uint64(1); i <= r.MessageIndex; i++ {
|
||||
@@ -68,54 +54,3 @@ func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
||||
func (cs *ConnectionState) Curve() cert.Curve {
|
||||
return cs.myCert.Curve()
|
||||
}
|
||||
|
||||
func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) ([]byte, error) {
|
||||
cs.decryptLock.Lock()
|
||||
result := cs.window.Check(l, messageCounter)
|
||||
cs.decryptLock.Unlock()
|
||||
if !result {
|
||||
return nil, ErrAlreadySeen
|
||||
}
|
||||
|
||||
out, err := cs.dKey.DecryptDanger(packet[header.Len:header.Len], packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
cs.decryptLock.Lock()
|
||||
result = cs.window.Update(l, messageCounter)
|
||||
cs.decryptLock.Unlock()
|
||||
if !result {
|
||||
return nil, ErrAlreadySeen
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) error {
|
||||
cs.decryptLock.Lock()
|
||||
result := cs.window.Check(l, messageCounter)
|
||||
cs.decryptLock.Unlock()
|
||||
if !result {
|
||||
return ErrAlreadySeen
|
||||
}
|
||||
|
||||
// The entire body is sent as AD, not encrypted.
|
||||
// The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value.
|
||||
// The packet is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's
|
||||
// otherwise malformed (meaning, there is no trailing 16 byte AEAD value), then this will result in at worst a 0-length slice
|
||||
// which will gracefully fail in the DecryptDanger call.
|
||||
signedPayload := packet[:len(packet)-cs.dKey.Overhead()]
|
||||
signatureValue := packet[len(packet)-cs.dKey.Overhead():]
|
||||
_, err := cs.dKey.DecryptDanger(nil, signedPayload, signatureValue, messageCounter, nb)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
cs.decryptLock.Lock()
|
||||
result = cs.window.Update(l, messageCounter)
|
||||
cs.decryptLock.Unlock()
|
||||
if !result {
|
||||
return ErrAlreadySeen
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
+1
-9
@@ -53,7 +53,6 @@ type Control struct {
|
||||
statsStart func()
|
||||
dnsStart func()
|
||||
lighthouseStart func()
|
||||
networkChangeStart func(rebind func())
|
||||
connectionManagerStart func(context.Context)
|
||||
}
|
||||
|
||||
@@ -105,9 +104,6 @@ func (c *Control) Start() error {
|
||||
if c.dnsStart != nil {
|
||||
go c.dnsStart()
|
||||
}
|
||||
if c.networkChangeStart != nil {
|
||||
go c.networkChangeStart(c.RebindUDPServer)
|
||||
}
|
||||
if c.connectionManagerStart != nil {
|
||||
go c.connectionManagerStart(c.ctx)
|
||||
}
|
||||
@@ -202,11 +198,7 @@ func (c *Control) RebindUDPServer() {
|
||||
return
|
||||
}
|
||||
|
||||
// A failure here means we are likely still pinned to the interface we came up on, so the rest of this is
|
||||
// unlikely to help. Say so instead of silently carrying on as if we rebound.
|
||||
if err := c.f.outside.Rebind(); err != nil {
|
||||
c.l.Error("Failed to rebind udp socket", "error", err)
|
||||
}
|
||||
_ = c.f.outside.Rebind()
|
||||
|
||||
// Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0
|
||||
c.f.lightHouse.SendUpdate()
|
||||
|
||||
@@ -78,7 +78,7 @@ func newReadyControl(t *testing.T) (*Control, *fakeDevice, *fakeConn) {
|
||||
inside: dev,
|
||||
outside: conn,
|
||||
writers: []udp.Conn{conn},
|
||||
batchers: make([]*batch.MultiCoalescer, 1),
|
||||
batchers: make([]batch.RxBatcher, 1),
|
||||
routines: 1,
|
||||
hostMap: newHostMap(l),
|
||||
lightHouse: lh,
|
||||
@@ -148,8 +148,8 @@ func (c *fakeConn) Rebind() error { c.rebinds++; ret
|
||||
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
||||
func (c *fakeConn) ListenOut(_ udp.EncReader, _ func()) error { return nil }
|
||||
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
|
||||
func (c *fakeConn) WriteBatch(bufs [][]byte, _ []netip.AddrPort) (int, error) {
|
||||
return len(bufs), nil
|
||||
func (c *fakeConn) WriteBatch(_ [][]byte, _ []netip.AddrPort, _ []byte) error {
|
||||
return nil
|
||||
}
|
||||
func (c *fakeConn) ReloadConfig(_ *config.C) {}
|
||||
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
|
||||
@@ -177,7 +177,7 @@ func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
||||
inside: dev,
|
||||
outside: conn,
|
||||
writers: []udp.Conn{conn},
|
||||
batchers: make([]*batch.MultiCoalescer, 2),
|
||||
batchers: make([]batch.RxBatcher, 2),
|
||||
routines: 2,
|
||||
l: test.NewLogger(),
|
||||
}
|
||||
|
||||
+1
-13
@@ -108,19 +108,7 @@ func (c *Control) GetVpnAddrs() []netip.Addr {
|
||||
}
|
||||
|
||||
func (c *Control) GetUDPAddr() netip.AddrPort {
|
||||
return c.f.outside.(*udp.TesterConn).GetAddr()
|
||||
}
|
||||
|
||||
// SetUDPAddr moves this node to a new underlay address, standing in for a laptop waking up on a different
|
||||
// network. Register the new address with the router as well or nothing will route back.
|
||||
func (c *Control) SetUDPAddr(addr netip.AddrPort) {
|
||||
c.f.outside.(*udp.TesterConn).SetAddr(addr)
|
||||
}
|
||||
|
||||
// SetLocalAddrsFn replaces underlay address discovery so a test can advertise its simulated address instead of
|
||||
// whatever this machine's NICs happen to be. Call it before Start, SendUpdate reads it from the update worker.
|
||||
func (c *Control) SetLocalAddrsFn(fn func(*LocalAllowList) []netip.Addr) {
|
||||
c.f.lightHouse.localAddrsFn = fn
|
||||
return c.f.outside.(*udp.TesterConn).Addr
|
||||
}
|
||||
|
||||
func (c *Control) KillPendingTunnel(vpnIp netip.Addr) bool {
|
||||
|
||||
@@ -1,187 +0,0 @@
|
||||
// Package cpupick chooses which CPUs the tun reader threads pin to when the
|
||||
// operator has not chosen for us (tun.cpu_affinity). The stock spread —
|
||||
// allowed[i] for routine i — has two failure modes this package exists to fix:
|
||||
//
|
||||
// - every co-located nebula starts its spread at allowed[0], so N instances
|
||||
// on one box stack their readers onto the same cores, and allowed[0] is
|
||||
// usually CPU 0, the core housekeeping and default IRQ affinity already
|
||||
// favor;
|
||||
// - on heterogeneous CPUs (ARM big.LITTLE, Intel P/E hybrids, AMD compact
|
||||
// cores) low IDs are not necessarily fast cores, and pinning an encrypt
|
||||
// thread to an efficiency core caps that queue's throughput.
|
||||
//
|
||||
// Default instead returns a preference-ordered pin list: the allowed set
|
||||
// filtered to performance cores (when the platform distinguishes them and
|
||||
// enough remain for every routine), confined to a single NUMA node and spread
|
||||
// across distinct physical cores when the topology permits, CPU 0's physical
|
||||
// core demoted to last resort, and the order rotated by a stable per-instance
|
||||
// key so co-located instances spread instead of stacking.
|
||||
package cpupick
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
|
||||
"github.com/slackhq/nebula/util"
|
||||
)
|
||||
|
||||
// topology is the slice of machine layout arrange consults: the NUMA node
|
||||
// and the physical core behind each candidate CPU, plus which core CPU 0
|
||||
// lives on (zeroCore, -1 when unknown — tracked separately because CPU 0's
|
||||
// SMT sibling deserves demotion even when CPU 0 itself isn't a candidate).
|
||||
// Probed from sysfs on Linux; flatTopology stands in when the platform can't
|
||||
// say, which turns every topology rule into a no-op rather than a wrong
|
||||
// answer.
|
||||
type topology struct {
|
||||
nodeOf map[int]int
|
||||
coreOf map[int]int
|
||||
zeroCore int
|
||||
}
|
||||
|
||||
// flatTopology places every CPU on node 0 and on a physical core of its own.
|
||||
func flatTopology(cpus []int) topology {
|
||||
t := topology{
|
||||
nodeOf: make(map[int]int, len(cpus)),
|
||||
coreOf: make(map[int]int, len(cpus)),
|
||||
zeroCore: -1,
|
||||
}
|
||||
for i, c := range cpus {
|
||||
t.nodeOf[c] = 0
|
||||
t.coreOf[c] = i
|
||||
if c == 0 {
|
||||
t.zeroCore = i
|
||||
}
|
||||
}
|
||||
return t
|
||||
}
|
||||
|
||||
// Default computes the pin order for `routines` tun readers. key is any
|
||||
// stable per-instance value; the bound UDP port is ideal — distinct across
|
||||
// co-located instances, stable across restarts so benchmark runs stay
|
||||
// comparable. Returns nil when there is nothing useful to say (no affinity
|
||||
// support on this platform, lookup failure); callers keep their existing
|
||||
// fallback spread.
|
||||
func Default(routines int, key uint64, l *slog.Logger) []int {
|
||||
allowed, err := util.AllowedCPUs()
|
||||
if err != nil || len(allowed) == 0 {
|
||||
return nil
|
||||
}
|
||||
perf, signal := perfCPUs(allowed)
|
||||
cands := pickCandidates(allowed, perf, routines)
|
||||
if len(cands) == 0 {
|
||||
return nil
|
||||
}
|
||||
if len(perf) < routines {
|
||||
signal = ""
|
||||
}
|
||||
cpus := arrange(cands, readTopology(cands), routines, splitmix64(key))
|
||||
if l != nil {
|
||||
l.Info("chose default pin CPUs for tun readers",
|
||||
"cpus", cpus[:min(routines, len(cpus))],
|
||||
"perfSignal", signal)
|
||||
}
|
||||
return cpus
|
||||
}
|
||||
|
||||
// pickCandidates applies the enough-for-everyone guard: a perf filter that
|
||||
// leaves fewer candidates than routines is discarded — giving every reader
|
||||
// its own (possibly slow) core beats stacking two readers on a fast one.
|
||||
func pickCandidates(allowed, perf []int, routines int) []int {
|
||||
if len(perf) < routines {
|
||||
return allowed
|
||||
}
|
||||
return perf
|
||||
}
|
||||
|
||||
// arrange turns the candidate set into the final pin order:
|
||||
//
|
||||
// 1. NUMA: when at least one node holds enough candidates for every
|
||||
// routine, confine to one such node, chosen by the instance hash. The
|
||||
// readers share hostmap and cipher state, so splitting one instance
|
||||
// across nodes taxes every packet — and co-located instances that hash
|
||||
// to different nodes stop competing entirely. When no node is big
|
||||
// enough, span nodes rather than stack readers.
|
||||
// 2. Rotate the preferred candidates by the hash so instances spread.
|
||||
// 3. SMT: emit one thread per physical core before any of their siblings —
|
||||
// two encrypt threads on one core split its execution units. Siblings
|
||||
// still follow for the routines > cores case.
|
||||
// 4. CPU 0's whole physical core goes last: housekeeping and default IRQ
|
||||
// noise on CPU 0 bleeds into its SMT sibling too. Within that tail the
|
||||
// sibling precedes CPU 0 itself, which only catches the bleed-through.
|
||||
//
|
||||
// The rotation happens before the SMT pass so each instance's one-per-core
|
||||
// walk also starts at a different core, and CPU 0's core is excluded from
|
||||
// the rotation so no hash value can put it back at the front.
|
||||
func arrange(cands []int, topo topology, routines int, h uint64) []int {
|
||||
byNode := map[int][]int{}
|
||||
var nodes []int
|
||||
for _, c := range cands {
|
||||
n := topo.nodeOf[c]
|
||||
if _, ok := byNode[n]; !ok {
|
||||
nodes = append(nodes, n)
|
||||
}
|
||||
byNode[n] = append(byNode[n], c)
|
||||
}
|
||||
var eligible []int
|
||||
for _, n := range nodes {
|
||||
if len(byNode[n]) >= routines {
|
||||
eligible = append(eligible, n)
|
||||
}
|
||||
}
|
||||
if len(eligible) > 0 {
|
||||
cands = byNode[eligible[int(h%uint64(len(eligible)))]]
|
||||
}
|
||||
|
||||
// Split off CPU 0's core: its siblings tail the list, CPU 0 tails them.
|
||||
preferred := make([]int, 0, len(cands))
|
||||
var zeroTail []int
|
||||
hasZero := false
|
||||
for _, c := range cands {
|
||||
switch {
|
||||
case c == 0:
|
||||
hasZero = true
|
||||
case topo.zeroCore >= 0 && topo.coreOf[c] == topo.zeroCore:
|
||||
zeroTail = append(zeroTail, c)
|
||||
default:
|
||||
preferred = append(preferred, c)
|
||||
}
|
||||
}
|
||||
if hasZero {
|
||||
zeroTail = append(zeroTail, 0)
|
||||
}
|
||||
if len(preferred) == 0 {
|
||||
return zeroTail // CPU 0's core is all we have
|
||||
}
|
||||
|
||||
// The node pick consumed the low hash bits; rotate by the high ones so
|
||||
// the two choices stay independent.
|
||||
off := int((h >> 32) % uint64(len(preferred)))
|
||||
rot := make([]int, 0, len(preferred))
|
||||
rot = append(rot, preferred[off:]...)
|
||||
rot = append(rot, preferred[:off]...)
|
||||
|
||||
seenCore := make(map[int]bool, len(rot))
|
||||
out := make([]int, 0, len(cands))
|
||||
var siblings []int
|
||||
for _, c := range rot {
|
||||
g := topo.coreOf[c]
|
||||
if seenCore[g] {
|
||||
siblings = append(siblings, c)
|
||||
continue
|
||||
}
|
||||
seenCore[g] = true
|
||||
out = append(out, c)
|
||||
}
|
||||
out = append(out, siblings...)
|
||||
out = append(out, zeroTail...)
|
||||
return out
|
||||
}
|
||||
|
||||
// splitmix64 decorrelates instance keys before the selection modulos: ports
|
||||
// on one box often share spacing (4242/4243, or round steps like +1000) that
|
||||
// raw key%len arithmetic would fold onto the same offset.
|
||||
func splitmix64(x uint64) uint64 {
|
||||
x += 0x9e3779b97f4a7c15
|
||||
x = (x ^ (x >> 30)) * 0xbf58476d1ce4e5b9
|
||||
x = (x ^ (x >> 27)) * 0x94d049bb133111eb
|
||||
return x ^ (x >> 31)
|
||||
}
|
||||
@@ -1,171 +0,0 @@
|
||||
package cpupick
|
||||
|
||||
import (
|
||||
"slices"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// pairTopo builds a topology where consecutive candidate pairs are SMT
|
||||
// siblings: (cpus[0],cpus[1]) share a core, (cpus[2],cpus[3]) the next, ...
|
||||
// All CPUs land on node 0.
|
||||
func pairTopo(cpus []int) topology {
|
||||
t := topology{
|
||||
nodeOf: make(map[int]int, len(cpus)),
|
||||
coreOf: make(map[int]int, len(cpus)),
|
||||
zeroCore: -1,
|
||||
}
|
||||
for i, c := range cpus {
|
||||
t.nodeOf[c] = 0
|
||||
t.coreOf[c] = i / 2
|
||||
if c == 0 {
|
||||
t.zeroCore = i / 2
|
||||
}
|
||||
}
|
||||
return t
|
||||
}
|
||||
|
||||
func TestArrangeDemotesZeroForEveryKey(t *testing.T) {
|
||||
candidates := []int{0, 1, 2, 3, 4, 5, 6, 7}
|
||||
for key := range uint64(64) {
|
||||
got := arrange(candidates, flatTopology(candidates), 4, splitmix64(key))
|
||||
if len(got) != len(candidates) {
|
||||
t.Fatalf("key %d: len=%d want %d", key, len(got), len(candidates))
|
||||
}
|
||||
if got[0] == 0 {
|
||||
t.Errorf("key %d: CPU 0 at the front: %v", key, got)
|
||||
}
|
||||
if got[len(got)-1] != 0 {
|
||||
t.Errorf("key %d: CPU 0 not demoted to last: %v", key, got)
|
||||
}
|
||||
sorted := slices.Clone(got)
|
||||
slices.Sort(sorted)
|
||||
if !slices.Equal(sorted, candidates) {
|
||||
t.Errorf("key %d: not a permutation: %v", key, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeDemotesZeroSiblings(t *testing.T) {
|
||||
// Pairs (0,1),(2,3),(4,5),(6,7): CPU 0's core — 0 and its sibling 1 —
|
||||
// must tail the list, sibling ahead of 0 itself.
|
||||
candidates := []int{0, 1, 2, 3, 4, 5, 6, 7}
|
||||
for key := range uint64(64) {
|
||||
got := arrange(candidates, pairTopo(candidates), 2, splitmix64(key))
|
||||
n := len(got)
|
||||
if got[n-1] != 0 || got[n-2] != 1 {
|
||||
t.Fatalf("key %d: tail = %v, want [... 1 0]", key, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeZeroSiblingWithoutZero(t *testing.T) {
|
||||
// CPU 0 excluded (cpuset) but its sibling 1 remains: the sibling still
|
||||
// tails the list when the topology knows which core CPU 0 lives on.
|
||||
candidates := []int{1, 2, 3, 4, 5}
|
||||
topo := pairTopo([]int{0, 1, 2, 3, 4, 5})
|
||||
got := arrange(candidates, topo, 2, splitmix64(7))
|
||||
if got[len(got)-1] != 1 {
|
||||
t.Errorf("CPU 0's sibling not demoted: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeRotatesByKey(t *testing.T) {
|
||||
candidates := []int{1, 2, 3, 4, 5, 6, 7, 8}
|
||||
seen := map[int]bool{}
|
||||
for key := range uint64(64) {
|
||||
seen[arrange(candidates, flatTopology(candidates), 4, splitmix64(key))[0]] = true
|
||||
}
|
||||
// 64 hashed keys over 8 slots must hit more than one starting CPU, or
|
||||
// co-located instances would all stack again.
|
||||
if len(seen) < 2 {
|
||||
t.Errorf("rotation never varied across keys: %v", seen)
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeStableForSameKey(t *testing.T) {
|
||||
candidates := []int{0, 2, 4, 6}
|
||||
topo := flatTopology(candidates)
|
||||
a := arrange(candidates, topo, 2, splitmix64(4242))
|
||||
b := arrange(candidates, topo, 2, splitmix64(4242))
|
||||
if !slices.Equal(a, b) {
|
||||
t.Errorf("same key ordered differently: %v vs %v", a, b)
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeZeroOnly(t *testing.T) {
|
||||
if got := arrange([]int{0}, flatTopology([]int{0}), 1, splitmix64(7)); !slices.Equal(got, []int{0}) {
|
||||
t.Errorf("sole CPU 0 must survive: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeSMTSiblingsLast(t *testing.T) {
|
||||
// Pairs (1,2),(3,4),(5,6),(7,8): the first four picks must cover four
|
||||
// distinct physical cores before any sibling repeats.
|
||||
candidates := []int{1, 2, 3, 4, 5, 6, 7, 8}
|
||||
topo := pairTopo(candidates)
|
||||
for key := range uint64(16) {
|
||||
got := arrange(candidates, topo, 4, splitmix64(key))
|
||||
seen := map[int]bool{}
|
||||
for _, c := range got[:4] {
|
||||
g := topo.coreOf[c]
|
||||
if seen[g] {
|
||||
t.Fatalf("key %d: sibling before all cores covered: %v", key, got)
|
||||
}
|
||||
seen[g] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeNUMAConfinesToOneNode(t *testing.T) {
|
||||
// Two nodes of four; both fit routines=3, so the result must sit
|
||||
// entirely inside one of them, and the hash must pick both across keys.
|
||||
candidates := []int{1, 2, 3, 4, 10, 11, 12, 13}
|
||||
topo := flatTopology(candidates)
|
||||
for _, c := range []int{10, 11, 12, 13} {
|
||||
topo.nodeOf[c] = 1
|
||||
}
|
||||
nodesSeen := map[int]bool{}
|
||||
for key := range uint64(32) {
|
||||
got := arrange(candidates, topo, 3, splitmix64(key))
|
||||
if len(got) != 4 {
|
||||
t.Fatalf("key %d: not confined to one node: %v", key, got)
|
||||
}
|
||||
n := topo.nodeOf[got[0]]
|
||||
for _, c := range got {
|
||||
if topo.nodeOf[c] != n {
|
||||
t.Fatalf("key %d: spans nodes: %v", key, got)
|
||||
}
|
||||
}
|
||||
nodesSeen[n] = true
|
||||
}
|
||||
if len(nodesSeen) != 2 {
|
||||
t.Errorf("hash never spread instances across nodes: %v", nodesSeen)
|
||||
}
|
||||
}
|
||||
|
||||
func TestArrangeNUMASpansWhenNoNodeFits(t *testing.T) {
|
||||
candidates := []int{1, 2, 3, 4, 10, 11, 12, 13}
|
||||
topo := flatTopology(candidates)
|
||||
for _, c := range []int{10, 11, 12, 13} {
|
||||
topo.nodeOf[c] = 1
|
||||
}
|
||||
got := arrange(candidates, topo, 6, splitmix64(1))
|
||||
if len(got) != len(candidates) {
|
||||
t.Errorf("undersized nodes must span, got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPickCandidates(t *testing.T) {
|
||||
allowed := []int{0, 1, 2, 3, 4, 5, 6, 7}
|
||||
perf := []int{4, 5}
|
||||
|
||||
// Enough perf cores for every routine: only they are used.
|
||||
if got := pickCandidates(allowed, perf, 2); !slices.Equal(got, perf) {
|
||||
t.Errorf("perf filter not applied: %v", got)
|
||||
}
|
||||
// Perf filter too small for the routine count: discarded, everyone
|
||||
// gets their own core from the full allowed set.
|
||||
if got := pickCandidates(allowed, perf, 4); !slices.Equal(got, allowed) {
|
||||
t.Errorf("undersized perf filter not discarded: %v", got)
|
||||
}
|
||||
}
|
||||
@@ -1,154 +0,0 @@
|
||||
//go:build linux
|
||||
|
||||
package cpupick
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// capacityKeepPct is the cpu_capacity admission threshold, relative to the
|
||||
// fastest allowed core. LITTLE cores are normalized to ~250-400 of the big
|
||||
// core's 1024 while mid cores sit at ~75%+, so half of max separates little
|
||||
// from the rest without splitting prime from mid on three-tier parts.
|
||||
const capacityKeepPct = 50
|
||||
|
||||
// freqKeepPct is the cpuinfo_max_freq admission threshold. Favored-core
|
||||
// turbo skew is 2-4% and ARM mid-vs-prime ~12%, while E-cores, LITTLE
|
||||
// cores, and AMD compact cores all sit >= 20% below their siblings' max.
|
||||
const freqKeepPct = 85
|
||||
|
||||
// perfCPUs partitions allowed into the subset that are "performance" cores,
|
||||
// consulting (in order of authority):
|
||||
//
|
||||
// 1. cpu_capacity — arch_topology's normalized per-CPU capacity, exposed on
|
||||
// arm/arm64/riscv; the scheduler's own view of big vs LITTLE.
|
||||
// 2. /sys/devices/cpu_core/cpus — the Intel hybrid P-core PMU mask, present
|
||||
// only on P/E parts (x86 has no cpu_capacity) and naming P cores outright.
|
||||
// 3. cpuinfo_max_freq — the cross-vendor fallback; catches AMD compact
|
||||
// cores, which neither of the above covers.
|
||||
//
|
||||
// Returns allowed unchanged (signal "") when nothing distinguishes the
|
||||
// cores: homogeneous parts, VMs without cpufreq, sysfs unavailable.
|
||||
func perfCPUs(allowed []int) ([]int, string) {
|
||||
return perfCPUsFrom("/sys/devices/system/cpu", "/sys/devices/cpu_core/cpus", allowed)
|
||||
}
|
||||
|
||||
func perfCPUsFrom(cpuDir, intelCoreMask string, allowed []int) ([]int, string) {
|
||||
if cpus, ok := byPerCPUValue(cpuDir, "cpu_capacity", allowed, capacityKeepPct); ok {
|
||||
return cpus, "cpu_capacity"
|
||||
}
|
||||
if cpus, ok := byIntelCoreMask(intelCoreMask, allowed); ok {
|
||||
return cpus, "intel_core_pmu"
|
||||
}
|
||||
if cpus, ok := byPerCPUValue(cpuDir, "cpufreq/cpuinfo_max_freq", allowed, freqKeepPct); ok {
|
||||
return cpus, "max_freq"
|
||||
}
|
||||
return allowed, ""
|
||||
}
|
||||
|
||||
// byPerCPUValue keeps the allowed CPUs whose per-CPU sysfs value is at least
|
||||
// keepPct percent of the maximum across allowed. Inconclusive (ok=false)
|
||||
// when any CPU is missing the file or when every value is equal.
|
||||
func byPerCPUValue(cpuDir, file string, allowed []int, keepPct int) ([]int, bool) {
|
||||
vals := make([]int, len(allowed))
|
||||
minV, maxV := 0, 0
|
||||
for i, cpu := range allowed {
|
||||
v, err := readIntFile(filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), file))
|
||||
if err != nil {
|
||||
return nil, false
|
||||
}
|
||||
vals[i] = v
|
||||
if i == 0 || v < minV {
|
||||
minV = v
|
||||
}
|
||||
if v > maxV {
|
||||
maxV = v
|
||||
}
|
||||
}
|
||||
if minV == maxV {
|
||||
return nil, false // homogeneous by this signal; try the next one
|
||||
}
|
||||
keep := make([]int, 0, len(allowed))
|
||||
for i, cpu := range allowed {
|
||||
if vals[i]*100 >= maxV*keepPct {
|
||||
keep = append(keep, cpu)
|
||||
}
|
||||
}
|
||||
return keep, true
|
||||
}
|
||||
|
||||
// byIntelCoreMask keeps the allowed CPUs named by the hybrid P-core PMU
|
||||
// mask. Inconclusive when the file is absent (non-hybrid x86, other arches)
|
||||
// or no allowed CPU is in the mask (the process was deliberately confined
|
||||
// to E-cores; nothing useful to prefer within that).
|
||||
func byIntelCoreMask(maskPath string, allowed []int) ([]int, bool) {
|
||||
b, err := os.ReadFile(maskPath)
|
||||
if err != nil {
|
||||
return nil, false
|
||||
}
|
||||
set, err := parseCPUList(strings.TrimSpace(string(b)))
|
||||
if err != nil || len(set) == 0 {
|
||||
return nil, false
|
||||
}
|
||||
pcore := make(map[int]bool, len(set))
|
||||
for _, c := range set {
|
||||
pcore[c] = true
|
||||
}
|
||||
keep := make([]int, 0, len(allowed))
|
||||
for _, cpu := range allowed {
|
||||
if pcore[cpu] {
|
||||
keep = append(keep, cpu)
|
||||
}
|
||||
}
|
||||
if len(keep) == 0 {
|
||||
return nil, false
|
||||
}
|
||||
return keep, true
|
||||
}
|
||||
|
||||
// parseCPUList decodes the kernel's cpulist format ("0-7,16-23", "3") into
|
||||
// individual CPU IDs. Empty input yields an empty list.
|
||||
func parseCPUList(s string) ([]int, error) {
|
||||
if s == "" {
|
||||
return nil, nil
|
||||
}
|
||||
var out []int
|
||||
for part := range strings.SplitSeq(s, ",") {
|
||||
part = strings.TrimSpace(part)
|
||||
if part == "" {
|
||||
continue
|
||||
}
|
||||
lo, hi, isRange := strings.Cut(part, "-")
|
||||
a, err := strconv.Atoi(lo)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bad cpulist entry %q: %w", part, err)
|
||||
}
|
||||
if !isRange {
|
||||
out = append(out, a)
|
||||
continue
|
||||
}
|
||||
b, err := strconv.Atoi(hi)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bad cpulist entry %q: %w", part, err)
|
||||
}
|
||||
if b < a || b-a > 8192 {
|
||||
return nil, fmt.Errorf("bad cpulist range %q", part)
|
||||
}
|
||||
for v := a; v <= b; v++ {
|
||||
out = append(out, v)
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func readIntFile(path string) (int, error) {
|
||||
b, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return strconv.Atoi(strings.TrimSpace(string(b)))
|
||||
}
|
||||
@@ -1,163 +0,0 @@
|
||||
//go:build linux
|
||||
|
||||
package cpupick
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// fakeSysfs builds a cpuDir tree with the given per-CPU file values.
|
||||
// A nil map for a file means "file absent on every CPU".
|
||||
func fakeSysfs(t *testing.T, capacity, maxFreq map[int]int) string {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
write := func(cpu int, rel string, v int) {
|
||||
p := filepath.Join(dir, fmt.Sprintf("cpu%d", cpu), rel)
|
||||
if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(p, fmt.Appendf(nil, "%d\n", v), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
for cpu, v := range capacity {
|
||||
write(cpu, "cpu_capacity", v)
|
||||
}
|
||||
for cpu, v := range maxFreq {
|
||||
write(cpu, "cpufreq/cpuinfo_max_freq", v)
|
||||
}
|
||||
return dir
|
||||
}
|
||||
|
||||
func writeCoreMask(t *testing.T, mask string) string {
|
||||
t.Helper()
|
||||
p := filepath.Join(t.TempDir(), "cpus")
|
||||
if err := os.WriteFile(p, []byte(mask+"\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
func TestPerfCPUsBigLittleCapacity(t *testing.T) {
|
||||
// 4 big (1024) + 4 LITTLE (~290): capacity is authoritative on ARM.
|
||||
dir := fakeSysfs(t, map[int]int{
|
||||
0: 1024, 1: 1024, 2: 1024, 3: 1024,
|
||||
4: 290, 5: 290, 6: 290, 7: 290,
|
||||
}, nil)
|
||||
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3, 4, 5, 6, 7})
|
||||
if signal != "cpu_capacity" {
|
||||
t.Fatalf("signal = %q", signal)
|
||||
}
|
||||
if !slices.Equal(got, []int{0, 1, 2, 3}) {
|
||||
t.Errorf("got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerfCPUsThreeTierKeepsMid(t *testing.T) {
|
||||
// prime (1024) + mid (~780) + little (~280): 50% keeps prime+mid.
|
||||
dir := fakeSysfs(t, map[int]int{
|
||||
0: 280, 1: 280, 2: 280, 3: 280,
|
||||
4: 780, 5: 780, 6: 780,
|
||||
7: 1024,
|
||||
}, nil)
|
||||
got, _ := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3, 4, 5, 6, 7})
|
||||
if !slices.Equal(got, []int{4, 5, 6, 7}) {
|
||||
t.Errorf("got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerfCPUsIntelHybridMask(t *testing.T) {
|
||||
// No cpu_capacity on x86; the P-core PMU mask decides.
|
||||
dir := fakeSysfs(t, nil, nil)
|
||||
mask := writeCoreMask(t, "0-7")
|
||||
got, signal := perfCPUsFrom(dir, mask, []int{0, 1, 2, 3, 8, 9, 10, 11})
|
||||
if signal != "intel_core_pmu" {
|
||||
t.Fatalf("signal = %q", signal)
|
||||
}
|
||||
if !slices.Equal(got, []int{0, 1, 2, 3}) {
|
||||
t.Errorf("got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerfCPUsIntelMaskDisjointFallsThrough(t *testing.T) {
|
||||
// Confined to E-cores only: the mask can't help, and equal freqs below
|
||||
// mean nothing else distinguishes them either -> allowed unchanged.
|
||||
dir := fakeSysfs(t, nil, map[int]int{8: 4300000, 9: 4300000})
|
||||
mask := writeCoreMask(t, "0-7")
|
||||
got, signal := perfCPUsFrom(dir, mask, []int{8, 9})
|
||||
if signal != "" || !slices.Equal(got, []int{8, 9}) {
|
||||
t.Errorf("got %v signal %q", got, signal)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerfCPUsMaxFreqCompactCores(t *testing.T) {
|
||||
// AMD-style compact cores: no capacity, no Intel mask; 3.3 vs 5.7 GHz.
|
||||
dir := fakeSysfs(t, nil, map[int]int{
|
||||
0: 5700000, 1: 5700000, 2: 3300000, 3: 3300000,
|
||||
})
|
||||
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3})
|
||||
if signal != "max_freq" {
|
||||
t.Fatalf("signal = %q", signal)
|
||||
}
|
||||
if !slices.Equal(got, []int{0, 1}) {
|
||||
t.Errorf("got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerfCPUsFavoredCoreSkewKept(t *testing.T) {
|
||||
// Turbo Boost Max favored cores run a few percent hot; they must not
|
||||
// shrink the candidate set to one or two cores.
|
||||
dir := fakeSysfs(t, nil, map[int]int{
|
||||
0: 5800000, 1: 5700000, 2: 5700000, 3: 5600000,
|
||||
})
|
||||
got, _ := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3})
|
||||
if !slices.Equal(got, []int{0, 1, 2, 3}) {
|
||||
t.Errorf("favored-core skew filtered CPUs: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerfCPUsHomogeneousInconclusive(t *testing.T) {
|
||||
dir := fakeSysfs(t, nil, map[int]int{0: 3000000, 1: 3000000})
|
||||
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1})
|
||||
if signal != "" || !slices.Equal(got, []int{0, 1}) {
|
||||
t.Errorf("got %v signal %q", got, signal)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerfCPUsNoSysfs(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2})
|
||||
if signal != "" || !slices.Equal(got, []int{0, 1, 2}) {
|
||||
t.Errorf("got %v signal %q", got, signal)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCPUList(t *testing.T) {
|
||||
cases := []struct {
|
||||
in string
|
||||
want []int
|
||||
wantErr bool
|
||||
}{
|
||||
{"0-3", []int{0, 1, 2, 3}, false},
|
||||
{"0-1,16-17", []int{0, 1, 16, 17}, false},
|
||||
{"5", []int{5}, false},
|
||||
{"", nil, false},
|
||||
{"3-1", nil, true},
|
||||
{"a-b", nil, true},
|
||||
{"1,x", nil, true},
|
||||
}
|
||||
for _, c := range cases {
|
||||
got, err := parseCPUList(c.in)
|
||||
if (err != nil) != c.wantErr {
|
||||
t.Errorf("%q: err=%v wantErr=%v", c.in, err, c.wantErr)
|
||||
continue
|
||||
}
|
||||
if !c.wantErr && !slices.Equal(got, c.want) {
|
||||
t.Errorf("%q: got %v want %v", c.in, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,10 +0,0 @@
|
||||
//go:build !linux
|
||||
|
||||
package cpupick
|
||||
|
||||
// perfCPUs is Linux-only sysfs walking; elsewhere report "no distinction".
|
||||
// Default already returns nil off-Linux (util.AllowedCPUs has no answer
|
||||
// there), so this exists to keep the package compiling everywhere.
|
||||
func perfCPUs(allowed []int) ([]int, string) {
|
||||
return allowed, ""
|
||||
}
|
||||
@@ -1,118 +0,0 @@
|
||||
//go:build linux
|
||||
|
||||
package cpupick
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// readTopology probes the NUMA node and physical-core layout of cpus from
|
||||
// sysfs. Anything sysfs won't say degrades toward flatTopology: an unknown
|
||||
// node becomes node 0, an unknown core becomes a core of its own — either
|
||||
// way the corresponding arrange rule becomes a no-op instead of a wrong
|
||||
// answer.
|
||||
func readTopology(cpus []int) topology {
|
||||
return readTopologyFrom("/sys/devices/system/node", "/sys/devices/system/cpu", cpus)
|
||||
}
|
||||
|
||||
func readTopologyFrom(nodeDir, cpuDir string, cpus []int) topology {
|
||||
coreOf, zeroCore := coreGroups(cpuDir, cpus)
|
||||
return topology{
|
||||
nodeOf: numaNodes(nodeDir, cpus),
|
||||
coreOf: coreOf,
|
||||
zeroCore: zeroCore,
|
||||
}
|
||||
}
|
||||
|
||||
// numaNodes maps each cpu to its NUMA node via
|
||||
// /sys/devices/system/node/nodeN/cpulist. CPUs no node claims (or no node
|
||||
// dirs at all: VMs, non-NUMA kernels) land on node 0.
|
||||
func numaNodes(nodeDir string, cpus []int) map[int]int {
|
||||
out := make(map[int]int, len(cpus))
|
||||
for _, c := range cpus {
|
||||
out[c] = 0
|
||||
}
|
||||
entries, err := os.ReadDir(nodeDir)
|
||||
if err != nil {
|
||||
return out
|
||||
}
|
||||
want := make(map[int]bool, len(cpus))
|
||||
for _, c := range cpus {
|
||||
want[c] = true
|
||||
}
|
||||
for _, e := range entries {
|
||||
id, ok := strings.CutPrefix(e.Name(), "node")
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
n, err := strconv.Atoi(id)
|
||||
if err != nil {
|
||||
continue // has_cpu, possible, ... share the prefix
|
||||
}
|
||||
b, err := os.ReadFile(filepath.Join(nodeDir, e.Name(), "cpulist"))
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
list, err := parseCPUList(strings.TrimSpace(string(b)))
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
for _, c := range list {
|
||||
if want[c] {
|
||||
out[c] = n
|
||||
}
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// coreGroups maps each cpu to a dense physical-core id derived from its
|
||||
// (physical_package_id, core_id) pair — core_id alone repeats across
|
||||
// sockets. CPUs whose topology files are unreadable get a core of their own.
|
||||
// The second return is the group id of the core CPU 0 lives on, or -1 when
|
||||
// that can't be determined; CPU 0's own files are consulted even when 0 is
|
||||
// not a candidate, so its SMT siblings are recognized under cpusets that
|
||||
// exclude CPU 0 itself.
|
||||
func coreGroups(cpuDir string, cpus []int) (map[int]int, int) {
|
||||
type pkgCore struct{ pkg, core int }
|
||||
pairOf := func(cpu int) (pkgCore, bool) {
|
||||
topoDir := filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), "topology")
|
||||
pkg, err1 := readIntFile(filepath.Join(topoDir, "physical_package_id"))
|
||||
core, err2 := readIntFile(filepath.Join(topoDir, "core_id"))
|
||||
if err1 != nil || err2 != nil {
|
||||
return pkgCore{}, false
|
||||
}
|
||||
return pkgCore{pkg, core}, true
|
||||
}
|
||||
|
||||
ids := map[pkgCore]int{}
|
||||
out := make(map[int]int, len(cpus))
|
||||
next := 0
|
||||
for _, cpu := range cpus {
|
||||
k, ok := pairOf(cpu)
|
||||
if !ok {
|
||||
out[cpu] = next
|
||||
next++
|
||||
continue
|
||||
}
|
||||
id, ok := ids[k]
|
||||
if !ok {
|
||||
id = next
|
||||
next++
|
||||
ids[k] = id
|
||||
}
|
||||
out[cpu] = id
|
||||
}
|
||||
|
||||
zeroCore := -1
|
||||
if k, ok := pairOf(0); ok {
|
||||
if id, ok := ids[k]; ok {
|
||||
zeroCore = id
|
||||
}
|
||||
}
|
||||
return out, zeroCore
|
||||
}
|
||||
@@ -1,111 +0,0 @@
|
||||
//go:build linux
|
||||
|
||||
package cpupick
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// fakeTopoSysfs builds nodeDir/cpuDir trees. nodes maps node id -> cpulist
|
||||
// string; cores maps cpu -> (package, core) pair.
|
||||
func fakeTopoSysfs(t *testing.T, nodes map[int]string, cores map[int][2]int) (string, string) {
|
||||
t.Helper()
|
||||
base := t.TempDir()
|
||||
nodeDir := filepath.Join(base, "node")
|
||||
cpuDir := filepath.Join(base, "cpu")
|
||||
for n, list := range nodes {
|
||||
d := filepath.Join(nodeDir, fmt.Sprintf("node%d", n))
|
||||
if err := os.MkdirAll(d, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(d, "cpulist"), []byte(list+"\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
for cpu, pc := range cores {
|
||||
d := filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), "topology")
|
||||
if err := os.MkdirAll(d, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(d, "physical_package_id"), fmt.Appendf(nil, "%d\n", pc[0]), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(d, "core_id"), fmt.Appendf(nil, "%d\n", pc[1]), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
return nodeDir, cpuDir
|
||||
}
|
||||
|
||||
func TestReadTopology(t *testing.T) {
|
||||
// Two nodes; SMT pairs (0,4),(1,5) on node 0 and (2,6),(3,7) on node 1.
|
||||
// core_id repeats across packages on purpose: the pair must disambiguate.
|
||||
nodeDir, cpuDir := fakeTopoSysfs(t,
|
||||
map[int]string{0: "0-1,4-5", 1: "2-3,6-7"},
|
||||
map[int][2]int{
|
||||
0: {0, 0}, 4: {0, 0}, 1: {0, 1}, 5: {0, 1},
|
||||
2: {1, 0}, 6: {1, 0}, 3: {1, 1}, 7: {1, 1},
|
||||
})
|
||||
cpus := []int{0, 1, 2, 3, 4, 5, 6, 7}
|
||||
topo := readTopologyFrom(nodeDir, cpuDir, cpus)
|
||||
|
||||
for _, c := range []int{0, 1, 4, 5} {
|
||||
if topo.nodeOf[c] != 0 {
|
||||
t.Errorf("cpu %d on node %d, want 0", c, topo.nodeOf[c])
|
||||
}
|
||||
}
|
||||
for _, c := range []int{2, 3, 6, 7} {
|
||||
if topo.nodeOf[c] != 1 {
|
||||
t.Errorf("cpu %d on node %d, want 1", c, topo.nodeOf[c])
|
||||
}
|
||||
}
|
||||
pairs := [][2]int{{0, 4}, {1, 5}, {2, 6}, {3, 7}}
|
||||
for _, p := range pairs {
|
||||
if topo.coreOf[p[0]] != topo.coreOf[p[1]] {
|
||||
t.Errorf("siblings %v not grouped: %d vs %d", p, topo.coreOf[p[0]], topo.coreOf[p[1]])
|
||||
}
|
||||
}
|
||||
if topo.coreOf[0] == topo.coreOf[2] {
|
||||
t.Error("cross-package cores with equal core_id must not merge")
|
||||
}
|
||||
if topo.zeroCore != topo.coreOf[0] {
|
||||
t.Errorf("zeroCore = %d, want %d", topo.zeroCore, topo.coreOf[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadTopologyZeroCoreWithoutZeroCandidate(t *testing.T) {
|
||||
// CPU 0 is not a candidate (cpuset excludes it) but its sibling 4 is:
|
||||
// zeroCore must still identify their shared core.
|
||||
nodeDir, cpuDir := fakeTopoSysfs(t,
|
||||
map[int]string{0: "0-7"},
|
||||
map[int][2]int{0: {0, 0}, 4: {0, 0}, 1: {0, 1}, 5: {0, 1}})
|
||||
topo := readTopologyFrom(nodeDir, cpuDir, []int{1, 4, 5})
|
||||
if topo.zeroCore < 0 || topo.coreOf[4] != topo.zeroCore {
|
||||
t.Errorf("zeroCore = %d, coreOf[4] = %d; sibling of CPU 0 not identified", topo.zeroCore, topo.coreOf[4])
|
||||
}
|
||||
if topo.coreOf[1] == topo.zeroCore {
|
||||
t.Error("cpu 1 wrongly grouped with CPU 0's core")
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadTopologyMissingSysfs(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
cpus := []int{0, 1, 2}
|
||||
topo := readTopologyFrom(filepath.Join(base, "nope"), filepath.Join(base, "also-nope"), cpus)
|
||||
seen := map[int]bool{}
|
||||
for _, c := range cpus {
|
||||
if topo.nodeOf[c] != 0 {
|
||||
t.Errorf("cpu %d node = %d, want 0", c, topo.nodeOf[c])
|
||||
}
|
||||
if seen[topo.coreOf[c]] {
|
||||
t.Errorf("cpu %d shares a fallback core group", c)
|
||||
}
|
||||
seen[topo.coreOf[c]] = true
|
||||
}
|
||||
if topo.zeroCore != -1 {
|
||||
t.Errorf("zeroCore = %d, want -1 when unknown", topo.zeroCore)
|
||||
}
|
||||
}
|
||||
@@ -1,10 +0,0 @@
|
||||
//go:build !linux
|
||||
|
||||
package cpupick
|
||||
|
||||
// readTopology has no sysfs to consult off Linux; the flat stand-in makes
|
||||
// arrange's NUMA and SMT rules no-ops. Default is already nil off Linux
|
||||
// (util.AllowedCPUs has no answer there) — this keeps the package compiling.
|
||||
func readTopology(cpus []int) topology {
|
||||
return flatTopology(cpus)
|
||||
}
|
||||
+7
-16
@@ -97,7 +97,8 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
||||
newAddr := getDnsServerAddr(c)
|
||||
|
||||
d.serverMu.Lock()
|
||||
running := d.server != nil
|
||||
running := d.server
|
||||
runningStarted := d.started
|
||||
sameAddr := d.addr == newAddr
|
||||
d.addr = newAddr
|
||||
d.enabled.Store(enabled)
|
||||
@@ -111,7 +112,7 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
||||
}
|
||||
|
||||
if !enabled {
|
||||
if running {
|
||||
if running != nil {
|
||||
d.Stop()
|
||||
}
|
||||
// Drop any records that accumulated while enabled; a later re-enable
|
||||
@@ -120,12 +121,12 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
if !running {
|
||||
if running == nil {
|
||||
// Was disabled (or never started); bring it up now.
|
||||
go d.Start()
|
||||
} else if !sameAddr {
|
||||
// Stop clears the slot before shutting down, otherwise the Start below can find the dying server and refuse
|
||||
d.Stop()
|
||||
d.shutdownServer(running, runningStarted, "reload")
|
||||
// Old Start goroutine has now exited; bring up a fresh listener on the new address.
|
||||
go d.Start()
|
||||
}
|
||||
|
||||
@@ -161,9 +162,7 @@ func (d *dnsServer) Start() {
|
||||
|
||||
started := make(chan struct{})
|
||||
d.serverMu.Lock()
|
||||
// Re-check enabled under the lock, a disable that raced our check above snapshots the slot under it too.
|
||||
// Two reloads in quick succession can both spawn a Start, the loser would orphan the live listener past Stop
|
||||
if d.ctx.Err() != nil || d.server != nil || !d.enabled.Load() {
|
||||
if d.ctx.Err() != nil {
|
||||
d.serverMu.Unlock()
|
||||
return
|
||||
}
|
||||
@@ -201,14 +200,6 @@ func (d *dnsServer) Start() {
|
||||
close(started)
|
||||
}
|
||||
|
||||
// Release our slot, unless a reload already replaced us, so a dead listener can't block a future Start
|
||||
d.serverMu.Lock()
|
||||
if d.server == server {
|
||||
d.server = nil
|
||||
d.started = nil
|
||||
}
|
||||
d.serverMu.Unlock()
|
||||
|
||||
if err != nil {
|
||||
d.l.Warn("Failed to run the DNS responder", "error", err)
|
||||
}
|
||||
|
||||
+4
-206
@@ -194,51 +194,14 @@ func TestDnsServer_reload_initial_serveDnsWithoutLighthouse(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDnsServer_reload_sameAddr_noOp(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
setDnsConfig(c, "127.0.0.1", "0", true, true)
|
||||
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
|
||||
go ds.Start()
|
||||
waitForBind(t, ds)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
before := ds.server
|
||||
ds.serverMu.Unlock()
|
||||
require.NotNil(t, before)
|
||||
|
||||
// Same address, so the running listener must be left alone rather than rebuilt under live queries
|
||||
// No server running yet, no addr change. Reload should not spawn anything.
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
assert.True(t, ds.enabled.Load())
|
||||
|
||||
ds.serverMu.Lock()
|
||||
after := ds.server
|
||||
ds.serverMu.Unlock()
|
||||
assert.Same(t, before, after, "a same-address reload must not restart the listener")
|
||||
|
||||
ds.Stop()
|
||||
}
|
||||
|
||||
// The branch the old sameAddr test was accidentally hitting: enabled with nothing running means reload starts it.
|
||||
func TestDnsServer_reload_whenNotRunning_starts(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
|
||||
// initial only records config, it never starts anything
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
ds.serverMu.Lock()
|
||||
assert.Nil(t, ds.server, "the initial reload must not start a listener")
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
waitForBind(t, ds)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
assert.NotNil(t, ds.server, "a reload with nothing running should bring DNS up")
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
ds.Stop()
|
||||
assert.Nil(t, ds.server)
|
||||
}
|
||||
|
||||
func TestDnsServer_StartStop_lifecycle(t *testing.T) {
|
||||
@@ -464,168 +427,3 @@ func waitFor(t *testing.T, cond func() bool) {
|
||||
}
|
||||
t.Fatal("timed out waiting for condition")
|
||||
}
|
||||
|
||||
// Two reloads in quick succession, or a HUP before Control.Start, can race two Starts at the same listener.
|
||||
func TestDnsServer_Start_isIdempotent(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
|
||||
go ds.Start()
|
||||
waitForBind(t, ds)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
first := ds.server
|
||||
ds.serverMu.Unlock()
|
||||
require.NotNil(t, first)
|
||||
|
||||
// If the second Start replaces the tracked server, Stop kills the wrong one and the port leaks
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
ds.Start()
|
||||
close(done)
|
||||
}()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("second Start never returned")
|
||||
}
|
||||
|
||||
ds.serverMu.Lock()
|
||||
second := ds.server
|
||||
ds.serverMu.Unlock()
|
||||
assert.Same(t, first, second, "a second Start must not replace the running server")
|
||||
|
||||
// The real proof, after Stop the port must actually be free
|
||||
ds.Stop()
|
||||
waitFor(t, func() bool {
|
||||
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
_ = pc.Close()
|
||||
return true
|
||||
})
|
||||
}
|
||||
|
||||
// An address change must actually end up listening on the new port. Start's guard refuses when a server is already
|
||||
// installed, so reload has to clear the slot before shutting the old one down.
|
||||
func TestDnsServer_reload_addrChange_restarts(t *testing.T) {
|
||||
first := freeUDPPort(t)
|
||||
second := freeUDPPort(t)
|
||||
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", first, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
|
||||
go ds.Start()
|
||||
waitForBind(t, ds)
|
||||
|
||||
// Cycle a few times, the failure this guards against depends on which goroutine wins serverMu
|
||||
for i := range 8 {
|
||||
want := second
|
||||
if i%2 == 1 {
|
||||
want = first
|
||||
}
|
||||
setDnsConfig(c, "127.0.0.1", want, true, true)
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
waitForBind(t, ds)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
srv := ds.server
|
||||
ds.serverMu.Unlock()
|
||||
require.NotNil(t, srv, "reload left DNS down instead of restarting it")
|
||||
require.Equal(t, "127.0.0.1:"+want, srv.Addr, "reload should be serving the new address")
|
||||
}
|
||||
|
||||
// Land back on second so the port assertions below are meaningful
|
||||
setDnsConfig(c, "127.0.0.1", second, true, true)
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
waitForBind(t, ds)
|
||||
|
||||
// The old port must be released and the new one actually held
|
||||
waitFor(t, func() bool {
|
||||
pc, err := net.ListenPacket("udp", "127.0.0.1:"+first)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
_ = pc.Close()
|
||||
return true
|
||||
})
|
||||
_, err := net.ListenPacket("udp", "127.0.0.1:"+second)
|
||||
require.Error(t, err, "the new address should be bound by the DNS responder")
|
||||
|
||||
ds.Stop()
|
||||
}
|
||||
|
||||
// A listener that dies on its own must release the slot, or a later same-addr reload sees it as running and no-ops.
|
||||
func TestDnsServer_Start_bindFailure_releasesSlot(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
blocker, err := net.ListenPacket("udp", "127.0.0.1:"+port)
|
||||
require.NoError(t, err)
|
||||
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
|
||||
ds.Start() // returns once the bind fails
|
||||
|
||||
ds.serverMu.Lock()
|
||||
assert.Nil(t, ds.server, "a listener that failed to bind must not stay parked in the slot")
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
// With the slot released, a reload can retry once the port frees up
|
||||
require.NoError(t, blocker.Close())
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
waitForBind(t, ds)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
assert.NotNil(t, ds.server, "a same-addr reload should retry after a failed bind")
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
ds.Stop()
|
||||
}
|
||||
|
||||
// A disable that lands while Start is between its unlocked check and the guard must not leave a listener behind.
|
||||
func TestDnsServer_Start_refusesWhenDisabledUnderLock(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
require.True(t, ds.enabled.Load())
|
||||
|
||||
// Holding serverMu parks Start on the lock, the only way to land the disable in that window on purpose
|
||||
ds.serverMu.Lock()
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
ds.Start()
|
||||
close(done)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
ds.serverMu.Unlock()
|
||||
t.Fatal("Start returned early, the test never exercised the window")
|
||||
case <-time.After(time.Millisecond * 100):
|
||||
}
|
||||
|
||||
// The disable reload's critical section. It sees nothing running, so it never calls Stop.
|
||||
ds.enabled.Store(false)
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("Start never returned")
|
||||
}
|
||||
|
||||
ds.serverMu.Lock()
|
||||
assert.Nil(t, ds.server, "Start must not install a listener a disable already cancelled")
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
|
||||
require.NoError(t, err, "an orphaned listener is still holding the port")
|
||||
_ = pc.Close()
|
||||
}
|
||||
|
||||
@@ -725,70 +725,6 @@ func TestReestablishRelays(t *testing.T) {
|
||||
|
||||
}
|
||||
|
||||
func TestRelayHandshakeOverDisestablishedEntry(t *testing.T) {
|
||||
t.Parallel()
|
||||
// If them tears down the tunnel while me keeps Established relay state, me's next
|
||||
// handshake flows through the relay with no fresh CreateRelayRequest and lands on
|
||||
// them's Disestablished terminal relay entry. them must re-establish that entry, or
|
||||
// its first transmit deletes its only relay and the tunnel is born transmit-dead:
|
||||
// them can receive but every send is silently dropped.
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them ", "10.128.0.2/24", m{"relay": m{"use_relays": true}})
|
||||
|
||||
// Teach my how to get to the relay and that their can be reached via the relay
|
||||
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
||||
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
||||
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||
|
||||
// Build a router so we don't have to reason who gets which packet
|
||||
r := router.NewR(t, myControl, relayControl, theirControl)
|
||||
defer r.RenderFlow()
|
||||
|
||||
// Start the servers
|
||||
myControl.Start()
|
||||
relayControl.Start()
|
||||
theirControl.Start()
|
||||
|
||||
t.Log("Trigger a handshake from me to them via the relay")
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
p := r.RouteForAllUntilTxTun(theirControl)
|
||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||
oldIdx := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false).LocalIndex
|
||||
|
||||
t.Log("Close the tunnel on them only, marking their relay entry Disestablished")
|
||||
theirControl.CloseTunnel(myVpnIpNet[0].Addr(), true)
|
||||
|
||||
t.Log("Re-handshake from me, riding the still-Established relay state")
|
||||
myControl.ReHandshake(theirVpnIpNet[0].Addr())
|
||||
for {
|
||||
h := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false)
|
||||
if h != nil && h.LocalIndex != oldIdx && h.RemoteIndex != 0 {
|
||||
break
|
||||
}
|
||||
r.RouteForAllExitFunc(func(*udp.Packet, *nebula.Control) router.ExitType {
|
||||
return router.RouteAndExit
|
||||
})
|
||||
}
|
||||
|
||||
hAtThem := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
|
||||
require.NotNil(t, hAtThem, "them should have completed the relayed handshake")
|
||||
require.Equal(t, []netip.Addr{relayVpnIpNet[0].Addr()}, hAtThem.CurrentRelaysToMe, "them should know a relay for the new tunnel")
|
||||
|
||||
t.Log("Send from them to me; their only relay entry must survive the transmit")
|
||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||
require.Never(t, func() bool {
|
||||
h := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
|
||||
return h == nil || len(h.CurrentRelaysToMe) == 0
|
||||
}, time.Second, 10*time.Millisecond, "them deleted its only relay entry; the tunnel is permanently transmit-dead")
|
||||
|
||||
p = r.RouteForAllUntilTxTun(myControl)
|
||||
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
|
||||
r.RenderHostmaps("Final hostmaps", myControl, relayControl, theirControl)
|
||||
}
|
||||
|
||||
func TestStage1RaceRelays(t *testing.T) {
|
||||
t.Parallel()
|
||||
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
||||
|
||||
@@ -1,225 +0,0 @@
|
||||
//go:build e2e_testing
|
||||
// +build e2e_testing
|
||||
|
||||
package e2e
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/e2e/router"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// reportedAddrs is what the lighthouse would hand a peer asking where vpnAddr is.
|
||||
func reportedAddrs(t *testing.T, lh *nebula.Control, vpnAddr netip.Addr) []netip.AddrPort {
|
||||
t.Helper()
|
||||
cm := lh.QueryLighthouse(vpnAddr)
|
||||
if cm == nil {
|
||||
return nil
|
||||
}
|
||||
var out []netip.AddrPort
|
||||
for _, c := range *cm {
|
||||
out = append(out, c.Reported...)
|
||||
out = append(out, c.Learned...)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// waitForLighthouseMsg routes until a lighthouse message lands on lh, or gives up. Reports whether one arrived.
|
||||
func waitForLighthouseMsg(t *testing.T, r *router.R, lh *nebula.Control, wait time.Duration) bool {
|
||||
t.Helper()
|
||||
h := &header.H{}
|
||||
return r.RouteForAllExitFuncOrTimeout(wait, func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||
if c != lh {
|
||||
return router.KeepRouting
|
||||
}
|
||||
// Punches are a single byte and never parse, they are just not what we are after
|
||||
if err := h.Parse(p.Data); err != nil {
|
||||
return router.KeepRouting
|
||||
}
|
||||
if h.Type == header.LightHouse {
|
||||
return router.RouteAndExit
|
||||
}
|
||||
return router.KeepRouting
|
||||
})
|
||||
}
|
||||
|
||||
// A laptop that changes networks has to tell the lighthouse promptly, otherwise the lighthouse keeps handing peers
|
||||
// the old address and their punches land nowhere. On a long lighthouse interval the only thing that closes that
|
||||
// window is the rebind, which on darwin the network change monitor drives. The e2e build compiles the monitor out,
|
||||
// so we call RebindUDPServer directly, which is the same thing the monitor does.
|
||||
func TestRebindSendsLighthouseUpdate(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
|
||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||
"lighthouse": m{"am_lighthouse": true},
|
||||
})
|
||||
|
||||
// 600s interval, so nothing scheduled can send an update during this test. A rebind is the only thing that can.
|
||||
myControl, _, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
|
||||
"lighthouse": m{
|
||||
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||
"interval": 600,
|
||||
},
|
||||
"static_host_map": m{
|
||||
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||
},
|
||||
})
|
||||
|
||||
r := router.NewR(t, lhControl, myControl)
|
||||
defer r.RenderFlow()
|
||||
|
||||
lhControl.Start()
|
||||
myControl.Start()
|
||||
|
||||
// Let the startup registration finish, then clear everything it left behind
|
||||
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
|
||||
r.RouteFor(time.Millisecond * 400)
|
||||
|
||||
// Nothing should be talking to the lighthouse on its own now
|
||||
require.False(t, waitForLighthouseMsg(t, r, lhControl, time.Millisecond*200),
|
||||
"nothing should reach the lighthouse before the rebind")
|
||||
|
||||
myControl.RebindUDPServer()
|
||||
|
||||
assert.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5),
|
||||
"a rebind should push an update to the lighthouse rather than waiting out the interval")
|
||||
|
||||
lhControl.Stop()
|
||||
myControl.Stop()
|
||||
}
|
||||
|
||||
// The other half of a rebind: every live tunnel requeries the lighthouse on its next send. That query is what makes
|
||||
// the lighthouse tell the peer to punch toward our new address, which is the part that actually revives a tunnel
|
||||
// whose remote NAT state died while we were on a different network.
|
||||
func TestRebindRequeriesPeersOnNextSend(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
|
||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||
"lighthouse": m{"am_lighthouse": true},
|
||||
})
|
||||
|
||||
lhCfg := m{
|
||||
"lighthouse": m{
|
||||
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||
"interval": 600,
|
||||
// Without this the peers advertise this machine's real addresses and then try to punch at them,
|
||||
// which the router has no route for.
|
||||
"local_allow_list": m{
|
||||
"10.0.0.0/24": true,
|
||||
"::/0": false,
|
||||
},
|
||||
},
|
||||
"static_host_map": m{
|
||||
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||
},
|
||||
}
|
||||
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", lhCfg)
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.3/24", lhCfg)
|
||||
|
||||
r := router.NewR(t, lhControl, myControl, theirControl)
|
||||
defer r.RenderFlow()
|
||||
|
||||
lhControl.Start()
|
||||
myControl.Start()
|
||||
theirControl.Start()
|
||||
r.RouteFor(time.Millisecond * 500)
|
||||
|
||||
// Point the peers at each other directly, this test is about the rebind and not about lighthouse discovery
|
||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
||||
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("initial")))
|
||||
r.RouteFor(time.Second)
|
||||
require.NotNil(t, myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false), "expected a tunnel to them")
|
||||
r.RouteFor(time.Millisecond * 300)
|
||||
|
||||
// Assert on what the peer sees rather than on lighthouse traffic. A query for them makes the lighthouse send
|
||||
// them a punch notification, which is the whole point. Our own update to the lighthouse sends them nothing,
|
||||
// so this cannot be satisfied by the update the rebind itself pushes.
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("quiet")))
|
||||
require.False(t, waitForLighthouseMsg(t, r, theirControl, time.Millisecond*300),
|
||||
"an ordinary send should not requery the lighthouse")
|
||||
|
||||
myControl.RebindUDPServer()
|
||||
r.RouteFor(time.Millisecond * 300) // let the update the rebind itself sends pass by
|
||||
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("after rebind")))
|
||||
assert.True(t, waitForLighthouseMsg(t, r, theirControl, time.Second*5),
|
||||
"the first send after a rebind should requery the lighthouse, which then tells the peer to punch at us")
|
||||
|
||||
lhControl.Stop()
|
||||
myControl.Stop()
|
||||
theirControl.Stop()
|
||||
}
|
||||
|
||||
// The scenario this whole thing exists for: a laptop sleeps at the office and wakes up at home on a new address.
|
||||
// Until it tells the lighthouse, the lighthouse keeps handing peers the office address, so their punches land
|
||||
// nowhere and the tunnel stays dead. On a long interval the rebind is the only thing that closes that window.
|
||||
func TestRebindAdvertisesNewAddressAfterMove(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
|
||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||
"lighthouse": m{"am_lighthouse": true},
|
||||
})
|
||||
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
|
||||
"lighthouse": m{
|
||||
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||
"interval": 600,
|
||||
},
|
||||
"static_host_map": m{
|
||||
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||
},
|
||||
})
|
||||
|
||||
// Advertise wherever we currently are rather than this machine's real NICs, read fresh each time so a move
|
||||
// is picked up.
|
||||
myControl.SetLocalAddrsFn(func(*nebula.LocalAllowList) []netip.Addr {
|
||||
return []netip.Addr{myControl.GetUDPAddr().Addr()}
|
||||
})
|
||||
|
||||
r := router.NewR(t, lhControl, myControl)
|
||||
defer r.RenderFlow()
|
||||
|
||||
lhControl.Start()
|
||||
myControl.Start()
|
||||
|
||||
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
|
||||
r.RouteFor(time.Millisecond * 400)
|
||||
|
||||
require.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), myUdpAddr,
|
||||
"the lighthouse should know the address we started on")
|
||||
|
||||
// Wake up somewhere else
|
||||
newAddr := netip.MustParseAddrPort("10.0.0.99:4242")
|
||||
myControl.SetUDPAddr(newAddr)
|
||||
r.AddRoute(newAddr.Addr(), newAddr.Port(), myControl)
|
||||
|
||||
// Nothing has told the lighthouse, and with interval 600 nothing scheduled will
|
||||
r.RouteFor(time.Millisecond * 400)
|
||||
require.NotContains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
|
||||
"the lighthouse should still be handing out the old address before the rebind")
|
||||
|
||||
myControl.RebindUDPServer()
|
||||
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an update after the rebind")
|
||||
r.RouteFor(time.Millisecond * 400)
|
||||
|
||||
assert.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
|
||||
"after the rebind the lighthouse should hand peers our new address")
|
||||
|
||||
lhControl.Stop()
|
||||
myControl.Stop()
|
||||
}
|
||||
@@ -1,136 +0,0 @@
|
||||
//go:build e2e_testing
|
||||
// +build e2e_testing
|
||||
|
||||
package e2e
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/e2e/router"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
)
|
||||
|
||||
// TestRecoveryTiming measures how long a tunnel takes to come back after the peer stops accepting our traffic,
|
||||
// which is what a laptop waking on a new network looks like from the peer's side: its NAT has no state for where
|
||||
// we are now, so everything we send disappears.
|
||||
//
|
||||
// It is a measurement, not a pass/fail assertion. Recovery is timed to the moment the peer punches back at us,
|
||||
// since that is when its NAT opens and the tunnel is usable again.
|
||||
//
|
||||
// go test -tags e2e_testing -v -run TestRecoveryTiming ./e2e/
|
||||
func TestRecoveryTiming(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
rebind bool
|
||||
}{
|
||||
{"no trigger", false},
|
||||
{"rebind counter", true},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
d, lost := measureRecovery(t, tc.rebind)
|
||||
t.Logf("RESULT %-16s recovered in %-9v (%d packets lost)", tc.name, d.Round(time.Millisecond), lost)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// measureRecovery returns how long until the peer punched back, and how many of our packets died meanwhile. When
|
||||
// rebind is true we call RebindUDPServer once the tunnel goes dark, which is what the darwin network change
|
||||
// monitor does and what iOS has always done. When false, nothing tells nebula anything is wrong.
|
||||
func measureRecovery(t *testing.T, rebind bool) (time.Duration, int) {
|
||||
t.Helper()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
|
||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||
"lighthouse": m{"am_lighthouse": true},
|
||||
})
|
||||
|
||||
peerCfg := m{
|
||||
"lighthouse": m{
|
||||
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||
"interval": 600,
|
||||
"local_allow_list": m{
|
||||
"10.0.0.0/24": true,
|
||||
"::/0": false,
|
||||
},
|
||||
},
|
||||
"static_host_map": m{
|
||||
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||
},
|
||||
}
|
||||
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", peerCfg)
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.3/24", peerCfg)
|
||||
|
||||
r := router.NewR(t, lhControl, myControl, theirControl)
|
||||
defer r.RenderFlow()
|
||||
defer func() {
|
||||
lhControl.Stop()
|
||||
myControl.Stop()
|
||||
theirControl.Stop()
|
||||
}()
|
||||
|
||||
lhControl.Start()
|
||||
myControl.Start()
|
||||
theirControl.Start()
|
||||
r.RouteFor(time.Millisecond * 500)
|
||||
|
||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
||||
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("establish")))
|
||||
r.RouteFor(time.Second)
|
||||
if myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false) == nil {
|
||||
t.Fatal("failed to establish the tunnel we are measuring")
|
||||
}
|
||||
r.RouteFor(time.Millisecond * 500)
|
||||
|
||||
// From here the peer's NAT has no state for us, everything we send it disappears
|
||||
start := time.Now()
|
||||
blackholed := 0
|
||||
var recovered time.Duration
|
||||
|
||||
if rebind {
|
||||
myControl.RebindUDPServer()
|
||||
}
|
||||
|
||||
// Keep the tun busy the way someone retrying a stalled connection would
|
||||
stop := make(chan struct{})
|
||||
defer close(stop)
|
||||
go func() {
|
||||
tick := time.NewTicker(time.Millisecond * 200)
|
||||
defer tick.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-stop:
|
||||
return
|
||||
case <-tick.C:
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(
|
||||
theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("retry")))
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
r.RouteForAllExitFuncOrTimeout(time.Second*30, func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||
if c == theirControl && p.From == myControl.GetUDPAddr() {
|
||||
blackholed++
|
||||
return router.Drop
|
||||
}
|
||||
|
||||
// The peer reaching us directly is the moment its NAT opened, whether that is a punch or a handshake
|
||||
if c == myControl && p.From == theirUdpAddr {
|
||||
recovered = time.Since(start)
|
||||
return router.RouteAndExit
|
||||
}
|
||||
|
||||
return router.KeepRouting
|
||||
})
|
||||
|
||||
if recovered == 0 {
|
||||
t.Fatalf("no recovery within 30s (%d packets blackholed)", blackholed)
|
||||
}
|
||||
return recovered, blackholed
|
||||
}
|
||||
+30
-149
@@ -114,28 +114,6 @@ type packet struct {
|
||||
packet *udp.Packet
|
||||
tun bool // a packet pulled off a tun device
|
||||
rx bool // the packet was received by a udp device
|
||||
|
||||
// h is the nebula header, parsed once when the packet is recorded. parseErr says why there isn't one, which
|
||||
// the flow log reports rather than hiding. Punchy sends a single byte, so an unparseable packet is normal.
|
||||
h header.H
|
||||
parseErr error
|
||||
}
|
||||
|
||||
// fromAddr and toAddr are the addresses this packet actually travelled between. Reading them off the control
|
||||
// instead would misreport the whole history once a test moves a node. Tun packets are synthesized without
|
||||
// addresses, so they fall back to the control.
|
||||
func (p *packet) fromAddr() netip.AddrPort {
|
||||
if p.tun || !p.packet.From.IsValid() {
|
||||
return p.from.GetUDPAddr()
|
||||
}
|
||||
return p.packet.From
|
||||
}
|
||||
|
||||
func (p *packet) toAddr() netip.AddrPort {
|
||||
if p.tun || !p.packet.To.IsValid() {
|
||||
return p.to.GetUDPAddr()
|
||||
}
|
||||
return p.packet.To
|
||||
}
|
||||
|
||||
func (p *packet) WasReceived() {
|
||||
@@ -153,9 +131,6 @@ const (
|
||||
ExitNow ExitType = 1
|
||||
// RouteAndExit routes this packet and exits immediately afterwards
|
||||
RouteAndExit ExitType = 2
|
||||
// Drop discards this packet without delivering it and keeps routing. Use it to simulate a blackhole, such as
|
||||
// a restrictive NAT refusing traffic from an address it has not seen.
|
||||
Drop ExitType = 3
|
||||
)
|
||||
|
||||
type ExitFunc func(packet *udp.Packet, receiver *nebula.Control) ExitType
|
||||
@@ -166,9 +141,7 @@ type ExitFunc func(packet *udp.Packet, receiver *nebula.Control) ExitType
|
||||
func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
// t.Name() contains a slash for subtests, so the flow log can land in a nested directory
|
||||
fn := filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name()))
|
||||
if err := os.MkdirAll(filepath.Dir(fn), 0755); err != nil {
|
||||
if err := os.MkdirAll("mermaid", 0755); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
@@ -179,7 +152,7 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
||||
outNat: make(map[outNatKey]netip.AddrPort),
|
||||
flow: []flowEntry{},
|
||||
ignoreFlows: []ignoreFlow{},
|
||||
fn: fn,
|
||||
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())),
|
||||
t: t,
|
||||
cancelRender: cancel,
|
||||
}
|
||||
@@ -276,7 +249,7 @@ func (r *R) renderFlow() {
|
||||
continue
|
||||
}
|
||||
|
||||
addr := e.packet.fromAddr()
|
||||
addr := e.packet.from.GetUDPAddr()
|
||||
if _, ok := participants[addr]; ok {
|
||||
continue
|
||||
}
|
||||
@@ -295,6 +268,7 @@ func (r *R) renderFlow() {
|
||||
}
|
||||
|
||||
// Print packets
|
||||
h := &header.H{}
|
||||
for _, e := range r.flow {
|
||||
if e.packet == nil {
|
||||
//fmt.Fprintf(f, " note over %s: %s\n", strings.Join(participantsVals, ", "), e.note)
|
||||
@@ -306,22 +280,21 @@ func (r *R) renderFlow() {
|
||||
fmt.Fprintln(f, r.formatUdpPacket(p))
|
||||
|
||||
} else {
|
||||
if err := h.Parse(p.packet.Data); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
line := "--x"
|
||||
if p.rx {
|
||||
line = "->>"
|
||||
}
|
||||
|
||||
detail := fmt.Sprintf("%s(%s), index %v, counter: %v",
|
||||
p.h.TypeName(), p.h.SubTypeName(), p.h.RemoteIndex, p.h.MessageCounter)
|
||||
if p.parseErr != nil {
|
||||
detail = fmt.Sprintf("unparsed, %v (%d bytes)", p.parseErr, len(p.packet.Data))
|
||||
}
|
||||
|
||||
fmt.Fprintf(f, " %s%s%s: %s\n",
|
||||
normalizeName(p.fromAddr().String()),
|
||||
fmt.Fprintf(f,
|
||||
" %s%s%s: %s(%s), index %v, counter: %v\n",
|
||||
normalizeName(p.from.GetUDPAddr().String()),
|
||||
line,
|
||||
normalizeName(p.toAddr().String()),
|
||||
detail,
|
||||
normalizeName(p.to.GetUDPAddr().String()),
|
||||
h.TypeName(), h.SubTypeName(), h.RemoteIndex, h.MessageCounter,
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -435,34 +408,29 @@ func (r *R) unlockedInjectFlow(from, to *nebula.Control, p *udp.Packet, tun bool
|
||||
|
||||
r.renderHostmaps(fmt.Sprintf("Packet %v", len(r.flow)))
|
||||
|
||||
var h header.H
|
||||
var parseErr error
|
||||
if !tun {
|
||||
parseErr = h.Parse(p.Data)
|
||||
}
|
||||
|
||||
// Decide before copying, the copy comes from a freelist and an ignored packet would never be released
|
||||
for _, i := range r.ignoreFlows {
|
||||
if tun {
|
||||
if i.tun.HasValue && i.tun.IsTrue {
|
||||
return nil
|
||||
}
|
||||
continue
|
||||
if len(r.ignoreFlows) > 0 {
|
||||
var h header.H
|
||||
err := h.Parse(p.Data)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
// A packet we could not parse has no type to match against, so no rule can ignore it
|
||||
if parseErr == nil && i.messageType == h.Type && i.subType == h.Subtype {
|
||||
return nil
|
||||
for _, i := range r.ignoreFlows {
|
||||
if !tun {
|
||||
if i.messageType == h.Type && i.subType == h.Subtype {
|
||||
return nil
|
||||
}
|
||||
} else if i.tun.HasValue && i.tun.IsTrue {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fp := &packet{
|
||||
from: from,
|
||||
to: to,
|
||||
packet: p.Copy(),
|
||||
tun: tun,
|
||||
h: h,
|
||||
parseErr: parseErr,
|
||||
from: from,
|
||||
to: to,
|
||||
packet: p.Copy(),
|
||||
tun: tun,
|
||||
}
|
||||
|
||||
r.flow = append(r.flow, flowEntry{packet: fp})
|
||||
@@ -692,10 +660,6 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
||||
p.Release()
|
||||
return
|
||||
|
||||
case Drop:
|
||||
// Record it so the flow log shows the attempt, but never hand it to the receiver
|
||||
r.unlockedInjectFlow(sender, receiver, p, false)
|
||||
|
||||
case KeepRouting:
|
||||
fp := r.unlockedInjectFlow(sender, receiver, p, false)
|
||||
receiver.InjectUDPPacket(p)
|
||||
@@ -726,85 +690,6 @@ func (r *R) RouteUntilAfterMsgType(sender *nebula.Control, msgType header.Messag
|
||||
})
|
||||
}
|
||||
|
||||
// RouteFor routes everything that shows up for the given duration and then returns. Use it to let a test settle
|
||||
// deterministically rather than sleeping and hoping: a single FlushAll races a completing handshake, which queues
|
||||
// more packets right behind it.
|
||||
func (r *R) RouteFor(d time.Duration) {
|
||||
r.RouteForAllExitFuncOrTimeout(d, func(*udp.Packet, *nebula.Control) ExitType {
|
||||
return KeepRouting
|
||||
})
|
||||
}
|
||||
|
||||
// RouteForAllExitFuncOrTimeout is RouteForAllExitFunc with a deadline, reporting whether whatDo asked to exit
|
||||
// before time ran out. The unbounded version blocks forever on a quiet network, so this is what a test needs to
|
||||
// assert that something does NOT happen, or to route for a fixed settling period.
|
||||
func (r *R) RouteForAllExitFuncOrTimeout(timeout time.Duration, whatDo ExitFunc) bool {
|
||||
sc := make([]reflect.SelectCase, 0, len(r.controls)+1)
|
||||
cm := make([]*nebula.Control, 0, len(r.controls))
|
||||
|
||||
for _, c := range r.controls {
|
||||
sc = append(sc, reflect.SelectCase{
|
||||
Dir: reflect.SelectRecv,
|
||||
Chan: reflect.ValueOf(c.GetUDPTxChan()),
|
||||
Send: reflect.Value{},
|
||||
})
|
||||
cm = append(cm, c)
|
||||
}
|
||||
|
||||
timer := time.NewTimer(timeout)
|
||||
defer timer.Stop()
|
||||
sc = append(sc, reflect.SelectCase{
|
||||
Dir: reflect.SelectRecv,
|
||||
Chan: reflect.ValueOf(timer.C),
|
||||
Send: reflect.Value{},
|
||||
})
|
||||
|
||||
for {
|
||||
x, rx, _ := reflect.Select(sc)
|
||||
if x == len(cm) {
|
||||
return false
|
||||
}
|
||||
|
||||
r.Lock()
|
||||
p := rx.Interface().(*udp.Packet)
|
||||
receiver := r.getControl(cm[x].GetUDPAddr(), p.To, p)
|
||||
if receiver == nil {
|
||||
r.Unlock()
|
||||
panic("Can't RouteForAllExitFuncOrTimeout for host: " + p.To.String())
|
||||
}
|
||||
|
||||
e := whatDo(p, receiver)
|
||||
switch e {
|
||||
case ExitNow:
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
return true
|
||||
|
||||
case RouteAndExit:
|
||||
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||
receiver.InjectUDPPacket(p)
|
||||
fp.WasReceived()
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
return true
|
||||
|
||||
case Drop:
|
||||
// Record it so the flow log shows the attempt, but never hand it to the receiver
|
||||
r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||
|
||||
case KeepRouting:
|
||||
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||
receiver.InjectUDPPacket(p)
|
||||
fp.WasReceived()
|
||||
|
||||
default:
|
||||
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
|
||||
}
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
}
|
||||
}
|
||||
|
||||
func (r *R) RouteForAllUntilAfterMsgTypeTo(receiver *nebula.Control, msgType header.MessageType, subType header.MessageSubType) {
|
||||
h := &header.H{}
|
||||
r.RouteForAllExitFunc(func(p *udp.Packet, r *nebula.Control) ExitType {
|
||||
@@ -897,10 +782,6 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
||||
p.Release()
|
||||
return
|
||||
|
||||
case Drop:
|
||||
// Record it so the flow log shows the attempt, but never hand it to the receiver
|
||||
r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||
|
||||
case KeepRouting:
|
||||
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||
receiver.InjectUDPPacket(p)
|
||||
|
||||
@@ -0,0 +1,188 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"log/slog"
|
||||
"testing"
|
||||
|
||||
"golang.org/x/net/ipv4"
|
||||
)
|
||||
|
||||
func TestInnerECN(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
pkt []byte
|
||||
want byte
|
||||
}{
|
||||
{"empty", nil, 0},
|
||||
{"v4_NotECT", v4WithToS(0x00), 0x00},
|
||||
{"v4_ECT0", v4WithToS(0x02), 0x02},
|
||||
{"v4_ECT1", v4WithToS(0x01), 0x01},
|
||||
{"v4_CE", v4WithToS(0x03), 0x03},
|
||||
{"v4_DSCP_then_NotECT", v4WithToS(0x88 | 0x00), 0x00},
|
||||
{"v4_DSCP_then_CE", v4WithToS(0x88 | 0x03), 0x03},
|
||||
{"v6_NotECT", v6WithTC(0x00), 0x00},
|
||||
{"v6_ECT0", v6WithTC(0x02), 0x02},
|
||||
{"v6_CE", v6WithTC(0x03), 0x03},
|
||||
{"v6_DSCP_then_CE", v6WithTC(0x88 | 0x03), 0x03},
|
||||
{"unknown_version", []byte{0xa5, 0xff}, 0},
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
got := innerECN(c.pkt)
|
||||
if got != c.want {
|
||||
t.Errorf("innerECN=0x%02x want 0x%02x", got, c.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// v4WithToS returns a 2-byte slice tall enough for innerECN: byte 0 carries
|
||||
// version=4 in the high nibble, byte 1 is the full ToS so we exercise both
|
||||
// the DSCP and ECN portions through the byte 1 mask.
|
||||
func v4WithToS(tos byte) []byte {
|
||||
return []byte{0x45, tos}
|
||||
}
|
||||
|
||||
// v6WithTC builds a 2-byte slice that places a known traffic class value
|
||||
// across bytes 0 (high nibble of TC) and 1 (low nibble of TC). innerECN
|
||||
// extracts ECN as (b[1]>>4)&0x03, which corresponds to TC[1:0].
|
||||
func v6WithTC(tc byte) []byte {
|
||||
return []byte{0x60 | (tc>>4)&0x0f, (tc & 0x0f) << 4}
|
||||
}
|
||||
|
||||
func TestApplyOuterECN(t *testing.T) {
|
||||
silent := slog.New(slog.DiscardHandler)
|
||||
hi := &HostInfo{}
|
||||
|
||||
// Build a v4 packet helper with a given inner ECN field.
|
||||
v4 := func(innerECN byte) []byte {
|
||||
// 20-byte minimal IPv4 header with ToS = innerECN (DSCP zeroed).
|
||||
return []byte{
|
||||
0x45, innerECN, 0, 28,
|
||||
0, 0, 0x40, 0,
|
||||
64, 6, 0, 0,
|
||||
10, 0, 0, 1,
|
||||
10, 0, 0, 2,
|
||||
}
|
||||
}
|
||||
// Build a v6 packet helper with a given inner ECN field. ECN occupies
|
||||
// TC[1:0] which sit at byte 1 mask 0x30.
|
||||
v6 := func(innerECN byte) []byte {
|
||||
// 40-byte minimal IPv6 header with TC[1:0] = innerECN.
|
||||
pkt := make([]byte, 40)
|
||||
pkt[0] = 0x60 // version=6, TC[7:4]=0
|
||||
pkt[1] = (innerECN & 0x03) << 4 // TC[3:0]: low 2 bits = ECN, top 2 = DSCP-low (0)
|
||||
return pkt
|
||||
}
|
||||
|
||||
type cell struct {
|
||||
outer byte
|
||||
inner byte
|
||||
wantECN byte
|
||||
wantSame bool // expect inner unchanged (true => verify the byte didn't move)
|
||||
}
|
||||
|
||||
// RFC 6040 normal-mode combine table. Only outer==CE causes mutation.
|
||||
table := []cell{
|
||||
{ecnNotECT, ecnNotECT, ecnNotECT, true},
|
||||
{ecnNotECT, ecnECT0, ecnECT0, true},
|
||||
{ecnNotECT, ecnECT1, ecnECT1, true},
|
||||
{ecnNotECT, ecnCE, ecnCE, true},
|
||||
|
||||
{ecnECT0, ecnNotECT, ecnNotECT, true},
|
||||
{ecnECT0, ecnECT0, ecnECT0, true},
|
||||
{ecnECT0, ecnECT1, ecnECT1, true},
|
||||
{ecnECT0, ecnCE, ecnCE, true},
|
||||
|
||||
{ecnECT1, ecnNotECT, ecnNotECT, true},
|
||||
{ecnECT1, ecnECT0, ecnECT0, true},
|
||||
{ecnECT1, ecnECT1, ecnECT1, true},
|
||||
{ecnECT1, ecnCE, ecnCE, true},
|
||||
|
||||
{ecnCE, ecnNotECT, ecnNotECT, true}, // legacy: log, leave alone
|
||||
{ecnCE, ecnECT0, ecnCE, false}, // CE folded in
|
||||
{ecnCE, ecnECT1, ecnCE, false},
|
||||
{ecnCE, ecnCE, ecnCE, true},
|
||||
}
|
||||
|
||||
for _, c := range table {
|
||||
t.Run("v4", func(t *testing.T) {
|
||||
pkt := v4(c.inner)
|
||||
applyOuterECN(pkt, c.outer, hi, silent)
|
||||
got := pkt[1] & 0x03
|
||||
if got != c.wantECN {
|
||||
t.Errorf("v4 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
|
||||
}
|
||||
})
|
||||
t.Run("v6", func(t *testing.T) {
|
||||
pkt := v6(c.inner)
|
||||
applyOuterECN(pkt, c.outer, hi, silent)
|
||||
got := (pkt[1] >> 4) & 0x03
|
||||
if got != c.wantECN {
|
||||
t.Errorf("v6 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestApplyOuterECN_IPv4ChecksumStaysValid guards against H1: folding an outer
|
||||
// CE mark into the inner IPv4 ToS byte must keep the IPv4 header checksum valid.
|
||||
// The passthrough emit paths write the packet verbatim, so a stale checksum
|
||||
// turns an underlay congestion mark into packet loss at the receiver.
|
||||
func TestApplyOuterECN_IPv4ChecksumStaysValid(t *testing.T) {
|
||||
silent := slog.New(slog.DiscardHandler)
|
||||
hi := &HostInfo{}
|
||||
|
||||
// 20-byte IPv4 header with DSCP=0x88 and inner ECN = ECT(0). Folding CE
|
||||
// flips only the low two bits of the ToS byte while leaving DSCP intact.
|
||||
pkt := []byte{
|
||||
0x45, 0x88 | ecnECT0, 0, 40,
|
||||
0x1c, 0x46, 0x40, 0x00,
|
||||
64, 6, 0, 0,
|
||||
10, 0, 0, 1,
|
||||
10, 0, 0, 2,
|
||||
}
|
||||
// Stamp a correct header checksum before the fold.
|
||||
binary.BigEndian.PutUint16(pkt[10:12], ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]))
|
||||
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
|
||||
t.Fatal("test setup: initial header checksum invalid")
|
||||
}
|
||||
|
||||
applyOuterECN(pkt, ecnCE, hi, silent)
|
||||
|
||||
// CE folded in, DSCP preserved.
|
||||
if got, want := pkt[1], byte(0x88|ecnCE); got != want {
|
||||
t.Fatalf("ToS after fold = 0x%02x, want 0x%02x", got, want)
|
||||
}
|
||||
// The incremental RFC 1624 update must leave the checksum valid and equal
|
||||
// to a full recompute over the mutated header.
|
||||
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
|
||||
t.Fatalf("IPv4 header checksum invalid after CE fold: 0x%04x", binary.BigEndian.Uint16(pkt[10:12]))
|
||||
}
|
||||
if got, want := binary.BigEndian.Uint16(pkt[10:12]), ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]); got != want {
|
||||
t.Fatalf("checksum = 0x%04x, full recompute = 0x%04x", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// ipv4HeaderChecksum computes the RFC 1071 IPv4 header checksum over hdr,
|
||||
// treating the checksum field (bytes 10:12) as zero.
|
||||
func ipv4HeaderChecksum(hdr []byte) uint16 {
|
||||
var sum uint32
|
||||
for i := 0; i+1 < len(hdr); i += 2 {
|
||||
if i == 10 {
|
||||
continue // checksum field
|
||||
}
|
||||
sum += uint32(hdr[i])<<8 | uint32(hdr[i+1])
|
||||
}
|
||||
for sum > 0xffff {
|
||||
sum = (sum >> 16) + (sum & 0xffff)
|
||||
}
|
||||
return ^uint16(sum)
|
||||
}
|
||||
|
||||
// ipv4HeaderChecksumValid reports whether the stored checksum matches a fresh
|
||||
// computation over the header.
|
||||
func ipv4HeaderChecksumValid(hdr []byte) bool {
|
||||
return binary.BigEndian.Uint16(hdr[10:12]) == ipv4HeaderChecksum(hdr)
|
||||
}
|
||||
+20
-18
@@ -131,9 +131,6 @@ listen:
|
||||
port: 4242
|
||||
# Sets the max number of packets to pull from the kernel for each syscall (under systems that support recvmmsg)
|
||||
# default is 64, does not support reload
|
||||
# Note: on Linux with UDP GRO (kernel 5.10+), each receive slot is sized for a full 64KiB coalesced
|
||||
# superpacket, so the receive scratch is batch * 64KiB per listening socket (~4MiB per routine at the
|
||||
# default of 64). Lower this to trade peak per-syscall throughput for memory on constrained hosts.
|
||||
#batch: 64
|
||||
# Configure socket buffers for the udp side (outside), leave unset to use the system defaults. Values will be doubled by the kernel
|
||||
# Default is net.core.rmem_default and net.core.wmem_default (/proc/sys/net/core/rmem_default and /proc/sys/net/core/rmem_default)
|
||||
@@ -149,14 +146,6 @@ listen:
|
||||
# Default true; set to false to leave WDF in charge of inbound decisions on the listener port. Not reloadable.
|
||||
#windows_bypass_wdf: true
|
||||
|
||||
# On macOS only
|
||||
# macOS scopes the udp socket to the interface it was created on, so moving between networks (wifi to wired,
|
||||
# office to home) leaves Nebula sending out an interface that no longer has a route. When true, Nebula watches
|
||||
# the routing socket and rebinds the listener once the change settles.
|
||||
# iOS does not use this, the host app drives the same rebind itself.
|
||||
# Default true. Not reloadable.
|
||||
#rebind_on_network_change: true
|
||||
|
||||
# By default, Nebula replies to packets it has no tunnel for with a "recv_error" packet. This packet helps speed up reconnection
|
||||
# in the case that Nebula on either side did not shut down cleanly. This response can be abused as a way to discover if Nebula is running
|
||||
# on a host though. This option lets you configure if you want to send "recv_error" packets always, never, or only to private network remotes.
|
||||
@@ -267,19 +256,22 @@ tun:
|
||||
|
||||
# Linux only. pin_threads pins each tun reader/encrypt OS thread to a single CPU. This keeps every goroutine's
|
||||
# batched sends flowing through one XPS-selected NIC TX ring, so packets within a flow stay ordered on the wire
|
||||
# instead of being sprayed across multiple TX rings and reordered. Not reloadable. Coerced to false if routines <= 1.
|
||||
# instead of being sprayed across multiple TX rings and reordered. Not reloadable.
|
||||
#
|
||||
# When cpu_affinity is unset, nebula picks CPUs that do NOT service any physical NIC's interrupts (read from
|
||||
# /sys/class/net/*/device/msi_irqs and /proc/irq/*/effective_affinity_list): an encrypt thread pinned onto a core
|
||||
# that also runs NAPI for a NIC RX queue fights the softirq for the core and collapses throughput for flows hashed
|
||||
# to that queue. If the NIC's vectors blanket every allowed CPU (many drivers default to one queue per core) the
|
||||
# avoidance logs and falls back to the old spread; narrow the NIC's queue/IRQ spread (e.g. `ethtool -X <dev>
|
||||
# equal N`) or set cpu_affinity explicitly to benefit.
|
||||
#pin_threads: true
|
||||
|
||||
# Linux only. cpu_affinity overrides which CPUs the tun reader threads pin to: a list of CPU IDs, one per routine
|
||||
# (see the top-level `routines` setting). Lists shorter than `routines` are modulo-cycled across the queues; extra
|
||||
# entries are ignored. IDs must be within the process's allowed CPU set, so this respects taskset / cgroup cpusets;
|
||||
# a non-integer or not-allowed entry disables the override and falls back to spreading queues across the allowed
|
||||
# CPUs. Only meaningful while pin_threads is true. Not reloadable.
|
||||
# When unset, the default spread prefers performance cores on heterogeneous CPUs (ARM big.LITTLE, Intel P/E
|
||||
# hybrids, AMD compact cores), keeps all readers on one NUMA node and on distinct physical cores when the
|
||||
# topology allows (SMT siblings last), leaves CPU 0's physical core as a last resort, and rotates its starting
|
||||
# point per instance (keyed by the bound UDP port) so co-located nebulas don't stack their readers onto the
|
||||
# same cores.
|
||||
# CPUs. Setting this disables the automatic NIC-IRQ avoidance described under pin_threads — prefer CPUs that don't
|
||||
# service your underlay NIC's RX queue IRQs. Only meaningful while pin_threads is true. Not reloadable.
|
||||
#cpu_affinity:
|
||||
# - 2
|
||||
# - 4
|
||||
@@ -420,6 +412,16 @@ logging:
|
||||
# This setting is reloadable
|
||||
#inactivity_timeout: 10m
|
||||
|
||||
# ecn (default true) propagates ECN (Explicit Congestion Notification) across the tunnel per RFC 6040: the inner
|
||||
# packet's ECN codepoint is copied onto the outer carrier header on encapsulation, and an outer CE ("congestion
|
||||
# experienced") mark is folded back into the inner header on decapsulation. On linux it additionally stamps
|
||||
# RTAX_FEATURE_ECN on the routes nebula installs, so the kernel actively negotiates ECN for connections to mesh
|
||||
# prefixes. Disable this only when an underlay middlebox mangles or clears ECN bits unpredictably.
|
||||
# This setting is reloadable, BUT flipping it at runtime only updates the datapath (the inner<->outer copy/combine).
|
||||
# The RTAX_FEATURE_ECN flag on already-installed routes is NOT revisited on reload, so nebula must be restarted for
|
||||
# the route half of this setting to take effect.
|
||||
#ecn: true
|
||||
|
||||
# Nebula security group configuration
|
||||
firewall:
|
||||
# Action to take when a packet is not allowed by the firewall rules.
|
||||
|
||||
@@ -8,15 +8,6 @@ Before=sshd.service
|
||||
Type=notify
|
||||
NotifyAccess=main
|
||||
SyslogIdentifier=nebula
|
||||
|
||||
# Uncomment to run as an unprivileged user with only CAP_NET_ADMIN. Requires a
|
||||
# nebula user that owns the config directory. Add CAP_NET_BIND_SERVICE to both
|
||||
# lines if any listener (lighthouse DNS, listen.port, stats, sshd) binds <1024.
|
||||
#User=nebula
|
||||
#Group=nebula
|
||||
#CapabilityBoundingSet=CAP_NET_ADMIN
|
||||
#AmbientCapabilities=CAP_NET_ADMIN
|
||||
|
||||
ExecReload=/bin/kill -HUP $MAINPID
|
||||
ExecStart=/usr/local/bin/nebula -config /etc/nebula/config.yml
|
||||
Restart=always
|
||||
|
||||
@@ -65,12 +65,3 @@ func (fp Packet) MarshalJSON() ([]byte, error) {
|
||||
"Fragment": fp.Fragment,
|
||||
})
|
||||
}
|
||||
|
||||
// ParsedPacket is a Packet plus the parse byproducts the RX path reuses
|
||||
type ParsedPacket struct {
|
||||
Packet
|
||||
IPHdrLen int
|
||||
// FragAny reports any fragmentation at all: MF flag or nonzero offset for IPv4, a fragment extension header for IPv6.
|
||||
// Distinct from Packet.Fragment, which is true only for NON-FIRST fragments
|
||||
FragAny bool
|
||||
}
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
module github.com/slackhq/nebula
|
||||
|
||||
go 1.26.0
|
||||
go 1.25.0
|
||||
|
||||
require (
|
||||
dario.cat/mergo v1.0.2
|
||||
@@ -24,12 +24,12 @@ require (
|
||||
github.com/vishvananda/netlink v1.3.1
|
||||
go.uber.org/goleak v1.3.0
|
||||
go.yaml.in/yaml/v3 v3.0.4
|
||||
golang.org/x/crypto v0.54.0
|
||||
golang.org/x/crypto v0.53.0
|
||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
||||
golang.org/x/net v0.57.0
|
||||
golang.org/x/sync v0.22.0
|
||||
golang.org/x/sys v0.47.0
|
||||
golang.org/x/term v0.45.0
|
||||
golang.org/x/net v0.56.0
|
||||
golang.org/x/sync v0.21.0
|
||||
golang.org/x/sys v0.46.0
|
||||
golang.org/x/term v0.44.0
|
||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
||||
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b
|
||||
golang.zx2c4.com/wireguard/windows v1.0.1
|
||||
|
||||
@@ -162,8 +162,8 @@ golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACk
|
||||
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
||||
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
||||
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
|
||||
golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw=
|
||||
golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk=
|
||||
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
|
||||
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
|
||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=
|
||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
|
||||
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
||||
@@ -182,8 +182,8 @@ golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLL
|
||||
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
|
||||
golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
|
||||
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
||||
golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE=
|
||||
golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU=
|
||||
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
|
||||
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
|
||||
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
||||
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
@@ -191,8 +191,8 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
|
||||
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
|
||||
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||
golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM=
|
||||
golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||
golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
@@ -208,11 +208,11 @@ golang.org/x/sys v0.0.0-20210124154548-22da62e12c0c/go.mod h1:h1NjWce9XRLGQEsW7w
|
||||
golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
|
||||
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
||||
golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
|
||||
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
|
||||
golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc=
|
||||
golang.org/x/term v0.44.0/go.mod h1:7ze4MdzUzLXpSAoFP1H0bOI9aXDqveSvatT5vKcFh2Y=
|
||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||
golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
|
||||
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||
|
||||
+5
-15
@@ -295,13 +295,7 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
||||
hm.messageMetrics.Tx(header.Handshake, hh.machine.Subtype(), 1)
|
||||
err := hm.outside.WriteTo(stage0, addr)
|
||||
if err != nil {
|
||||
// These repeat every attempt, so match the success log below and only shout when the remotes changed
|
||||
level := slog.LevelDebug
|
||||
if remotesHaveChanged {
|
||||
level = slog.LevelError
|
||||
}
|
||||
|
||||
hostinfo.logger(hm.l).Log(context.Background(), level, "Failed to send handshake message",
|
||||
hostinfo.logger(hm.l).Error("Failed to send handshake message",
|
||||
"udpAddr", addr,
|
||||
"initiatorIndex", hostinfo.localIndexId,
|
||||
"handshake", hsFields,
|
||||
@@ -535,9 +529,7 @@ func (hm *HandshakeManager) DeleteHostInfo(hostinfo *HostInfo) {
|
||||
|
||||
func (hm *HandshakeManager) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
||||
for _, addr := range hostinfo.vpnAddrs {
|
||||
if cur, ok := hm.vpnIps[addr]; ok && cur.hostinfo == hostinfo {
|
||||
delete(hm.vpnIps, addr)
|
||||
}
|
||||
delete(hm.vpnIps, addr)
|
||||
}
|
||||
|
||||
if len(hm.vpnIps) == 0 {
|
||||
@@ -975,9 +967,7 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
||||
nb := make([]byte, 12, 12)
|
||||
out := make([]byte, mtu)
|
||||
for _, cp := range hh.packetStore {
|
||||
// TODO: use a SendBatch here. Each callback lands in
|
||||
// sendNoMetrics -> WriteTo: one syscall per cached packet,
|
||||
// where one sendmmsg could flush the whole store.
|
||||
//todo use a sendbatcher
|
||||
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
||||
}
|
||||
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
||||
@@ -1088,8 +1078,8 @@ func (hm *HandshakeManager) sendHandshakeResponse(via ViaSender, msg []byte, hos
|
||||
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||
// We received a valid handshake on this relay, so make sure the relay
|
||||
// state reflects that, in case it had been marked Disestablished.
|
||||
via.relayHI.relayState.UpdateRelayForByIdxState(via.relay.LocalIndex, Established)
|
||||
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false, 0)
|
||||
via.relayHI.relayState.UpdateRelayForByIdxState(via.remoteIdx, Established)
|
||||
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
||||
f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -84,7 +84,7 @@ func (mw *mockEncWriter) SendMessageToVpnAddr(_ header.MessageType, _ header.Mes
|
||||
return
|
||||
}
|
||||
|
||||
func (mw *mockEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
|
||||
func (mw *mockEncWriter) SendVia(_ *HostInfo, _ *Relay, _, _, _ []byte, _ bool) {
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
+6
-11
@@ -190,18 +190,13 @@ func SubTypeName(t MessageType, s MessageSubType) string {
|
||||
}
|
||||
|
||||
func IsValidSubType(t MessageType, s MessageSubType) bool {
|
||||
switch t {
|
||||
case Message:
|
||||
return s == MessageNone || s == MessageRelay
|
||||
case Handshake:
|
||||
return s == HandshakeIXPSK0
|
||||
case Test:
|
||||
return s == TestReply || s == TestRequest
|
||||
case Control, CloseTunnel, RecvError, LightHouse:
|
||||
return s == 0
|
||||
default:
|
||||
return false
|
||||
if n, ok := subTypeMap[t]; ok {
|
||||
if _, ok := (*n)[s]; ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// NewHeader turns bytes into a header
|
||||
|
||||
@@ -102,57 +102,6 @@ func TestTypeMap(t *testing.T) {
|
||||
}, subTypeMap)
|
||||
}
|
||||
|
||||
// mapIsValidSubType is the pre-refactor, map-driven definition of a valid
|
||||
// subtype. IsValidSubType was reimplemented as an explicit switch; this keeps
|
||||
// the original behavior around so we can prove the switch is equivalent to it.
|
||||
func mapIsValidSubType(t MessageType, s MessageSubType) bool {
|
||||
if n, ok := subTypeMap[t]; ok {
|
||||
if _, ok := (*n)[s]; ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func TestIsValidSubType(t *testing.T) {
|
||||
// Explicit intent table: documents exactly which subtypes are valid so the
|
||||
// test stays meaningful even if both the switch and subTypeMap change.
|
||||
assert.True(t, IsValidSubType(Message, MessageNone))
|
||||
assert.True(t, IsValidSubType(Message, MessageRelay))
|
||||
assert.False(t, IsValidSubType(Message, 2))
|
||||
|
||||
assert.True(t, IsValidSubType(Handshake, HandshakeIXPSK0))
|
||||
// HandshakeXXPSK0 is defined but not a wire-valid subtype.
|
||||
assert.False(t, IsValidSubType(Handshake, HandshakeXXPSK0))
|
||||
|
||||
assert.True(t, IsValidSubType(Test, TestRequest))
|
||||
assert.True(t, IsValidSubType(Test, TestReply))
|
||||
assert.False(t, IsValidSubType(Test, 2))
|
||||
|
||||
// These types only ever carry subtype 0.
|
||||
for _, mt := range []MessageType{Control, CloseTunnel, RecvError, LightHouse} {
|
||||
assert.True(t, IsValidSubType(mt, 0), "type %d subtype 0 should be valid", mt)
|
||||
assert.False(t, IsValidSubType(mt, 1), "type %d subtype 1 should be invalid", mt)
|
||||
}
|
||||
|
||||
// Unknown/unassigned types are never valid.
|
||||
assert.False(t, IsValidSubType(99, 0))
|
||||
|
||||
// Exhaustive proof of equivalence with the original map-driven logic across
|
||||
// the entire (type, subtype) input space.
|
||||
for ti := 0; ti <= 0xff; ti++ {
|
||||
for si := 0; si <= 0xff; si++ {
|
||||
mt, mst := MessageType(ti), MessageSubType(si)
|
||||
assert.Equalf(t, mapIsValidSubType(mt, mst), IsValidSubType(mt, mst),
|
||||
"IsValidSubType(%d, %d) diverged from map-driven definition", ti, si)
|
||||
}
|
||||
}
|
||||
|
||||
// H method must delegate to the package function.
|
||||
assert.True(t, (&H{Type: Test, Subtype: TestReply}).IsValidSubType())
|
||||
assert.False(t, (&H{Type: Handshake, Subtype: HandshakeXXPSK0}).IsValidSubType())
|
||||
}
|
||||
|
||||
func TestHeader_String(t *testing.T) {
|
||||
assert.Equal(
|
||||
t,
|
||||
|
||||
+1
-11
@@ -287,6 +287,7 @@ type HostInfo struct {
|
||||
type ViaSender struct {
|
||||
UdpAddr netip.AddrPort
|
||||
relayHI *HostInfo // relayHI is the host info object of the relay
|
||||
remoteIdx uint32 // remoteIdx is the index included in the header of the received packet
|
||||
relay *Relay // relay contains the rest of the relay information, including the PeerIP of the host trying to communicate with us.
|
||||
IsRelayed bool // IsRelayed is true if the packet was sent through a relay
|
||||
}
|
||||
@@ -543,17 +544,6 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
|
||||
return final
|
||||
}
|
||||
|
||||
func (hm *HostMap) QueryIndexCached(index uint32, cache map[uint32]*HostInfo) *HostInfo {
|
||||
if out, ok := cache[index]; ok {
|
||||
return out
|
||||
}
|
||||
out := hm.QueryIndex(index)
|
||||
if out != nil {
|
||||
cache[index] = out
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func (hm *HostMap) QueryIndex(index uint32) *HostInfo {
|
||||
hm.RLock()
|
||||
if h, ok := hm.Indexes[index]; ok {
|
||||
|
||||
@@ -15,7 +15,7 @@ import (
|
||||
"github.com/slackhq/nebula/routing"
|
||||
)
|
||||
|
||||
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.ParsedPacket, nb []byte, sendBatch *batch.SendBatch, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
|
||||
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packet, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
|
||||
// borrowed: pkt.Bytes is owned by the originating tio.Queue and is
|
||||
// only valid until the next Read on that queue. Every consumer below
|
||||
// (parse, self-forward, handshake cache, sendInsideMessage) reads it
|
||||
@@ -74,7 +74,7 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Parse
|
||||
return
|
||||
}
|
||||
|
||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(&fwPacket.Packet, func(hh *HandshakeHostInfo) {
|
||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
||||
// borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt
|
||||
// bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket),
|
||||
// so retaining segments past the loop is safe.
|
||||
@@ -105,9 +105,9 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Parse
|
||||
return
|
||||
}
|
||||
|
||||
dropReason := f.firewall.Drop(fwPacket.Packet, false, hostinfo, f.pki.GetCAPool(), localCache)
|
||||
dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
|
||||
if dropReason == nil {
|
||||
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch)
|
||||
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch, rejectBuf, q)
|
||||
} else {
|
||||
f.rejectInside(packet, rejectBuf, q)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
@@ -126,6 +126,7 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
|
||||
c := ci.messageCounter.Add(1)
|
||||
|
||||
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
|
||||
f.connectionManager.Out(hostinfo)
|
||||
|
||||
out, encErr := ci.eKey.EncryptDanger(out, out, seg, c, nb)
|
||||
if noiseutil.EncryptLockNeeded {
|
||||
@@ -137,7 +138,8 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
|
||||
"udpAddr", hostinfo.GetRemote(),
|
||||
"counter", c,
|
||||
)
|
||||
// Skip this segment; the rest of the superpacket can still go out. TCP will retransmit anything we drop here.
|
||||
// Skip this segment; the rest of the superpacket can still
|
||||
// go out — TCP will retransmit anything we drop here.
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -149,19 +151,16 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
|
||||
// later sendmmsg flush. Segmentation is fused with encryption here so the
|
||||
// kernel-supplied superpacket bytes never get written into a separate
|
||||
// scratch arena: SegmentSuperpacket builds each segment's plaintext in
|
||||
// segScratch[:segLen] in turn, and we encrypt directly into a fresh SendBatch slot.
|
||||
func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []byte, sendBatch *batch.SendBatch) {
|
||||
// segScratch[:segLen] in turn, and we encrypt directly into a fresh
|
||||
// SendBatch slot.
|
||||
func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int) {
|
||||
ci := hostinfo.ConnectionState
|
||||
if ci.eKey == nil {
|
||||
return
|
||||
}
|
||||
|
||||
// One traffic-out mark covers every segment of the superpacket; doing it
|
||||
// per segment in sendInsideEncrypt paid an atomic store up to ~45 extra
|
||||
// times per TSO packet, inside writeLock under boring crypto.
|
||||
f.connectionManager.Out(hostinfo)
|
||||
|
||||
remote := hostinfo.GetRemote()
|
||||
ecnEnabled := f.ecnEnabled.Load()
|
||||
if hostinfo.lastRebindCount != f.rebindCount {
|
||||
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
|
||||
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
|
||||
@@ -212,7 +211,11 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
|
||||
return nil
|
||||
}
|
||||
|
||||
sendBatch.Commit(toSend, relayHostInfo.GetRemote())
|
||||
var ecn byte
|
||||
if ecnEnabled {
|
||||
ecn = innerECN(seg)
|
||||
}
|
||||
sendBatch.Commit(toSend, relayHostInfo.GetRemote(), ecn)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
@@ -230,14 +233,36 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
|
||||
return nil
|
||||
}
|
||||
|
||||
sendBatch.Commit(out, remote)
|
||||
var ecn byte
|
||||
if ecnEnabled {
|
||||
ecn = innerECN(seg)
|
||||
}
|
||||
sendBatch.Commit(out, remote, ecn)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).Error("Failed to segment superpacket for send", "error", err)
|
||||
hostinfo.logger(f.l).Error("Failed to segment superpacket for send",
|
||||
"error", err,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
// innerECN returns the 2-bit IP-level ECN codepoint of an inner IPv4 or IPv6
|
||||
// packet, or 0 if pkt is too short or its IP version is unrecognized. Used at
|
||||
// encap to copy the inner codepoint onto the outer carrier per RFC 6040.
|
||||
func innerECN(pkt []byte) byte {
|
||||
if len(pkt) < 2 {
|
||||
return 0
|
||||
}
|
||||
switch pkt[0] >> 4 {
|
||||
case 4:
|
||||
return pkt[1] & 0x03
|
||||
case 6:
|
||||
return (pkt[1] >> 4) & 0x03
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||
if !f.firewall.OutboundSendReject {
|
||||
return
|
||||
@@ -254,30 +279,27 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, rejectBuf []byte, q int) {
|
||||
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) {
|
||||
if !f.firewall.InboundSendReject {
|
||||
return
|
||||
}
|
||||
|
||||
// split rejectBuf to make sure we have room to write the plaintext rejection, then encrypt it, without trampling anything
|
||||
// we can't re-use packet, if we need to send an icmp reject, it won't be long enough.
|
||||
half := len(rejectBuf) / 2
|
||||
encryptBuf := rejectBuf[0:0:half] //the first half of rejectBuf's capacity, len set to 0
|
||||
buildBuf := rejectBuf[half:]
|
||||
|
||||
out := iputil.CreateRejectPacket(packet, buildBuf)
|
||||
out = iputil.CreateRejectPacket(packet, out)
|
||||
if len(out) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
if len(out) > iputil.MaxRejectPacketSize {
|
||||
if f.l.Enabled(context.Background(), slog.LevelInfo) {
|
||||
f.l.Info("rejectOutside: packet too big, not sending", "packet", packet, "outPacket", out)
|
||||
f.l.Info("rejectOutside: packet too big, not sending",
|
||||
"packet", packet,
|
||||
"outPacket", out,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, encryptBuf, q)
|
||||
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, packet, q)
|
||||
}
|
||||
|
||||
// Handshake will attempt to initiate a tunnel with the provided vpn address. This is a no-op if the tunnel is already established or being established
|
||||
@@ -371,7 +393,7 @@ func (f *Interface) getOrHandshakeConsiderRouting(fwPacket *firewall.Packet, cac
|
||||
}
|
||||
|
||||
func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) {
|
||||
fp := &firewall.ParsedPacket{}
|
||||
fp := &firewall.Packet{}
|
||||
err := newPacket(p, false, fp)
|
||||
if err != nil {
|
||||
f.l.Warn("error while parsing outgoing packet for firewall check", "error", err)
|
||||
@@ -379,7 +401,7 @@ func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubTyp
|
||||
}
|
||||
|
||||
// check if packet is in outbound fw rules
|
||||
dropReason := f.firewall.Drop(fp.Packet, false, hostinfo, f.pki.GetCAPool(), nil)
|
||||
dropReason := f.firewall.Drop(*fp, false, hostinfo, f.pki.GetCAPool(), nil)
|
||||
if dropReason != nil {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("dropping cached packet",
|
||||
@@ -492,14 +514,20 @@ func (f *Interface) prepareSendVia(via *HostInfo,
|
||||
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
||||
// out is a buffer used to store the result of the Encrypt operation
|
||||
// q indicates which writer to use to send the packet.
|
||||
func (f *Interface) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
|
||||
func (f *Interface) SendVia(via *HostInfo,
|
||||
relay *Relay,
|
||||
ad,
|
||||
nb,
|
||||
out []byte,
|
||||
nocopy bool,
|
||||
) {
|
||||
toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy)
|
||||
if err != nil {
|
||||
// already logged by prepareSendVia
|
||||
return
|
||||
}
|
||||
|
||||
err = f.writers[q].WriteTo(toSend, via.GetRemote())
|
||||
err = f.writers[0].WriteTo(toSend, via.GetRemote())
|
||||
if err != nil {
|
||||
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
||||
}
|
||||
@@ -573,7 +601,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
|
||||
"error", err,
|
||||
"udpAddr", hr,
|
||||
"udpAddr", remote,
|
||||
)
|
||||
}
|
||||
} else {
|
||||
@@ -588,7 +616,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
||||
)
|
||||
continue
|
||||
}
|
||||
f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true, q)
|
||||
f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true)
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
+79
-89
@@ -96,7 +96,12 @@ type Interface struct {
|
||||
// pinThreads controls whether listenIn pins each TUN reader OS thread to
|
||||
// a CPU at all (tun.pin_threads, default true). When false, threads are
|
||||
// left free to migrate as on stock nebula.
|
||||
pinThreads bool
|
||||
pinThreads bool
|
||||
// ecnEnabled gates RFC 6040 underlay ECN propagation. When true,
|
||||
// inside.go copies the inner ECN onto the outer carrier on encap and
|
||||
// decryptToTun folds outer CE into the inner header on decap. Toggle
|
||||
// via tunnels.ecn (default true).
|
||||
ecnEnabled atomic.Bool
|
||||
relayManager *relayManager
|
||||
|
||||
tryPromoteEvery atomic.Uint32
|
||||
@@ -115,12 +120,10 @@ type Interface struct {
|
||||
ctx context.Context
|
||||
writers []udp.Conn
|
||||
queues []tio.Queue
|
||||
// batchers is one per tun queue, wrapping queues[i]. readOutsidePackets
|
||||
// commits plaintext into the batcher; the plaintext is decrypted
|
||||
// in place inside the UDP receive buffers, so listenOut must call Flush
|
||||
// at the end of each UDP recvmmsg batch, before those buffers are
|
||||
// reused (every udp.Conn ListenOut guarantees that ordering).
|
||||
batchers []*batch.MultiCoalescer
|
||||
// batchers is one per tun queue, wrapping queues[i].
|
||||
// decryptToTun sends plaintext into the batch.RxBatcher;
|
||||
// listenOut calls its Flush at the end of each UDP recvmmsg batch.
|
||||
batchers []batch.RxBatcher
|
||||
wg sync.WaitGroup
|
||||
|
||||
// fatalErr holds the first unexpected reader error that caused shutdown.
|
||||
@@ -132,13 +135,18 @@ type Interface struct {
|
||||
metricHandshakes metrics.Histogram
|
||||
messageMetrics *MessageMetrics
|
||||
cachedPacketMetrics *cachedPacketMetrics
|
||||
metricTxDropped metrics.Counter
|
||||
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
type EncWriter interface {
|
||||
SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int)
|
||||
SendVia(via *HostInfo,
|
||||
relay *Relay,
|
||||
ad,
|
||||
nb,
|
||||
out []byte,
|
||||
nocopy bool,
|
||||
)
|
||||
SendMessageToVpnAddr(t header.MessageType, st header.MessageSubType, vpnAddr netip.Addr, p, nb, out []byte)
|
||||
SendMessageToHostInfo(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte)
|
||||
Handshake(vpnAddr netip.Addr)
|
||||
@@ -197,10 +205,6 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
||||
return nil, errors.New("no connection manager")
|
||||
}
|
||||
|
||||
if c.routines <= 1 {
|
||||
c.PinThreads = false //pinning is not useful unless there's more than one tun reader
|
||||
}
|
||||
|
||||
cs := c.pki.getCertState()
|
||||
ifce := &Interface{
|
||||
ctx: ctx,
|
||||
@@ -218,7 +222,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
||||
routines: c.routines,
|
||||
version: c.version,
|
||||
writers: make([]udp.Conn, c.routines),
|
||||
batchers: make([]*batch.MultiCoalescer, c.routines),
|
||||
batchers: make([]batch.RxBatcher, c.routines),
|
||||
myVpnNetworks: cs.myVpnNetworks,
|
||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||
myVpnAddrs: cs.myVpnAddrs,
|
||||
@@ -231,7 +235,6 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
||||
pinThreads: c.PinThreads,
|
||||
|
||||
metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)),
|
||||
metricTxDropped: metrics.GetOrRegisterCounter("udp.tx.dropped", nil),
|
||||
messageMetrics: c.MessageMetrics,
|
||||
cachedPacketMetrics: &cachedPacketMetrics{
|
||||
sent: metrics.GetOrRegisterCounter("hostinfo.cached_packets.sent", nil),
|
||||
@@ -285,13 +288,6 @@ func (f *Interface) activate() error {
|
||||
return err
|
||||
}
|
||||
if len(queues) < f.routines {
|
||||
// TODO: this clamp is only safe because it is unreachable when the
|
||||
// udp side has multiple readers (linux Queues opens exactly n or
|
||||
// errors; every other platform already clamped routines to 1 above).
|
||||
// If a platform ever returns fewer queues than routines with
|
||||
// SO_REUSEPORT sockets already bound, the surplus sockets get no
|
||||
// listenOut and the kernel blackholes every flow it hashes to them —
|
||||
// fail loudly or close the extra sockets instead.
|
||||
f.l.Warn("tun multiqueue is not supported on this platform, falling back to fewer routines",
|
||||
"requested", f.routines, "opened", len(queues))
|
||||
f.routines = len(queues)
|
||||
@@ -301,7 +297,18 @@ func (f *Interface) activate() error {
|
||||
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
||||
|
||||
for i := range f.queues {
|
||||
f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l)
|
||||
caps := tio.QueueCapabilities(f.queues[i])
|
||||
if caps.TSO || caps.USO {
|
||||
// Multi-lane: TCP gets coalesced when TSO is on, UDP when USO
|
||||
// is on, everything else (and either lane disabled) falls
|
||||
// through to passthrough so non-IP / non-TCP-UDP traffic still
|
||||
// reaches the TUN.
|
||||
arena := batch.NewArena(batch.DefaultMultiArenaCap)
|
||||
f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l, arena, caps.TSO, caps.USO)
|
||||
} else {
|
||||
arena := batch.NewArena(batch.DefaultPassthroughArenaCap)
|
||||
f.batchers[i] = batch.NewPassthrough(f.queues[i], arena.Reserve, arena.Reset)
|
||||
}
|
||||
}
|
||||
|
||||
// On error the caller owns the cleanup, Control.Start cancels the service context
|
||||
@@ -349,31 +356,6 @@ func (f *Interface) onFatal(err error) {
|
||||
}
|
||||
}
|
||||
|
||||
type rxContext struct {
|
||||
q int
|
||||
scratch []byte
|
||||
// nb is a re-usable nonce buffer for decrypt calls to use
|
||||
nb []byte
|
||||
h *header.H
|
||||
fwPacket *firewall.ParsedPacket
|
||||
hostmapCache map[uint32]*HostInfo
|
||||
lhh *LightHouseHandler
|
||||
ctCache *firewall.ConntrackCacheTicker
|
||||
}
|
||||
|
||||
func newRxContext(f *Interface, q int) *rxContext {
|
||||
return &rxContext{
|
||||
q: q,
|
||||
scratch: make([]byte, mtu),
|
||||
nb: make([]byte, 12, 12),
|
||||
h: &header.H{},
|
||||
fwPacket: &firewall.ParsedPacket{},
|
||||
hostmapCache: map[uint32]*HostInfo{},
|
||||
lhh: f.lightHouse.NewRequestHandler(),
|
||||
ctCache: firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout),
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) listenOut(i int) {
|
||||
var li udp.Conn
|
||||
if i > 0 {
|
||||
@@ -382,17 +364,21 @@ func (f *Interface) listenOut(i int) {
|
||||
li = f.outside
|
||||
}
|
||||
|
||||
rxc := newRxContext(f, i)
|
||||
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||
lhh := f.lightHouse.NewRequestHandler()
|
||||
h := &header.H{}
|
||||
fwPacket := &firewall.Packet{}
|
||||
nb := make([]byte, 12, 12)
|
||||
|
||||
listener := func(fromUdpAddr netip.AddrPort, payload []byte) {
|
||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, payload, rxc)
|
||||
listener := func(fromUdpAddr netip.AddrPort, payload []byte, meta udp.RxMeta) {
|
||||
plaintext := f.batchers[i].Reserve(len(payload))
|
||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get(), meta)
|
||||
}
|
||||
|
||||
flusher := func() {
|
||||
if err := f.batchers[i].Flush(); err != nil {
|
||||
f.l.Error("Failed to flush tun coalescer", "error", err)
|
||||
}
|
||||
clear(rxc.hostmapCache)
|
||||
}
|
||||
|
||||
err := li.ListenOut(listener, flusher)
|
||||
@@ -408,36 +394,32 @@ func (f *Interface) listenOut(i int) {
|
||||
f.l.Debug("underlay reader is done", "reader", i)
|
||||
}
|
||||
|
||||
func (f *Interface) pinThisThread(i int) {
|
||||
var cpu int
|
||||
if n := len(f.cpuAffinity); n > 0 {
|
||||
// Explicit tun.cpu_affinity list wins; parseCpuAffinity already
|
||||
// validated the entries against the allowed CPU set.
|
||||
cpu = f.cpuAffinity[i%n]
|
||||
} else if allowed, err := util.AllowedCPUs(); err == nil && len(allowed) > 0 {
|
||||
// Default: spread queues across the CPUs we're actually allowed to
|
||||
// run on. Under a cpuset/taskset mask these aren't 0..NumCPU-1, so
|
||||
// i % NumCPU would pick unrunnable IDs and every pin would fail.
|
||||
cpu = allowed[i%len(allowed)]
|
||||
} else {
|
||||
cpu = i % runtime.NumCPU()
|
||||
}
|
||||
if err := util.PinThreadToCPU(cpu); err != nil {
|
||||
f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) listenIn(queue tio.Queue, i int) {
|
||||
// Pinning this thread (and goroutine) to a single CPU keeps every sendmmsg from this goroutine going through the
|
||||
// same TX ring on the nic, so the wire sees per-flow order. Skip entirely when tun.pin_threads is false.
|
||||
if f.pinThreads {
|
||||
f.pinThisThread(i)
|
||||
var cpu int
|
||||
if n := len(f.cpuAffinity); n > 0 {
|
||||
// Explicit tun.cpu_affinity list wins; parseCpuAffinity already
|
||||
// validated the entries against the allowed CPU set.
|
||||
cpu = f.cpuAffinity[i%n]
|
||||
} else if allowed, err := util.AllowedCPUs(); err == nil && len(allowed) > 0 {
|
||||
// Default: spread queues across the CPUs we're actually allowed to
|
||||
// run on. Under a cpuset/taskset mask these aren't 0..NumCPU-1, so
|
||||
// i % NumCPU would pick unrunnable IDs and every pin would fail.
|
||||
cpu = allowed[i%len(allowed)]
|
||||
} else {
|
||||
cpu = i % runtime.NumCPU()
|
||||
}
|
||||
if err := util.PinThreadToCPU(cpu); err != nil {
|
||||
f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err)
|
||||
}
|
||||
}
|
||||
|
||||
rejectBuf := make([]byte, mtu)
|
||||
arenaSize := batch.SendBatchCap * (udp.MTU + 32)
|
||||
sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, arenaSize)
|
||||
fwPacket := &firewall.ParsedPacket{}
|
||||
fwPacket := &firewall.Packet{}
|
||||
nb := make([]byte, 12, 12)
|
||||
|
||||
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||
@@ -459,35 +441,26 @@ func (f *Interface) listenIn(queue tio.Queue, i int) {
|
||||
// accumulated so the first packets of a deep read drain
|
||||
// hit the wire while the rest are still being encrypted.
|
||||
if sb.Len() >= batch.SendBatchCap {
|
||||
f.flushSendBatch(sb, i)
|
||||
if err := sb.Flush(); err != nil {
|
||||
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
f.flushSendBatch(sb, i)
|
||||
if err := sb.Flush(); err != nil {
|
||||
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
|
||||
}
|
||||
}
|
||||
|
||||
f.l.Debug("overlay reader is done", "reader", i)
|
||||
}
|
||||
|
||||
// flushSendBatch drains sb to the underlay and accounts for anything it could not deliver. A shortfall means
|
||||
// specific destinations were undeliverable (a stale remote, a reject rule), which the backend logs per peer at
|
||||
// debug; here it is only a counter, so one unreachable peer cannot spam a log line per batch.
|
||||
func (f *Interface) flushSendBatch(sb *batch.SendBatch, q int) {
|
||||
queued := sb.Len()
|
||||
written, err := sb.Flush()
|
||||
if err != nil {
|
||||
f.l.Error("Failed to write outgoing batch", "error", err, "writer", q)
|
||||
}
|
||||
if dropped := queued - written; dropped > 0 {
|
||||
f.metricTxDropped.Inc(int64(dropped))
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
||||
c.RegisterReloadCallback(f.reloadFirewall)
|
||||
c.RegisterReloadCallback(f.reloadSendRecvError)
|
||||
c.RegisterReloadCallback(f.reloadAcceptRecvError)
|
||||
c.RegisterReloadCallback(f.reloadDisconnectInvalid)
|
||||
c.RegisterReloadCallback(f.reloadMisc)
|
||||
c.RegisterReloadCallback(f.reloadEcn)
|
||||
|
||||
for _, udpConn := range f.writers {
|
||||
c.RegisterReloadCallback(udpConn.ReloadConfig)
|
||||
@@ -620,6 +593,23 @@ func (f *Interface) reloadMisc(c *config.C) {
|
||||
}
|
||||
}
|
||||
|
||||
// reloadEcn syncs Interface.ecnEnabled with the tunnels.ecn config knob.
|
||||
// Default is enabled (RFC 6040 normal mode); set false on the rare path
|
||||
// where an underlay middlebox rewrites or drops ECN bits unpredictably.
|
||||
func (f *Interface) reloadEcn(c *config.C) {
|
||||
initial := c.InitialLoad()
|
||||
if initial || c.HasChanged("tunnels.ecn") {
|
||||
v := c.GetBool("tunnels.ecn", true)
|
||||
changed := f.ecnEnabled.Swap(v) != v
|
||||
if !initial {
|
||||
f.l.Info("tunnels.ecn changed", "enabled", v)
|
||||
if changed {
|
||||
f.l.Warn("tunnels.ecn datapath toggled, but route-level ECN negotiation (RTAX_FEATURE_ECN) retains its previous state until nebula is restarted", "enabled", v)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
||||
ticker := time.NewTicker(i)
|
||||
defer ticker.Stop()
|
||||
|
||||
+3
-27
@@ -199,7 +199,7 @@ func ipv4CreateRejectTCPPacket(packet []byte, out []byte) []byte {
|
||||
}
|
||||
|
||||
func ipv6CreateRejectPacket(packet []byte, out []byte) []byte {
|
||||
proto, offset, isFragment := IPv6FindUpperProtocol(packet)
|
||||
proto, offset, isFragment := ipv6FindUpperProtocol(packet)
|
||||
if isFragment {
|
||||
return nil
|
||||
}
|
||||
@@ -333,34 +333,11 @@ func ipv6CreateRejectTCPPacket(packet []byte, out []byte, offset int) []byte {
|
||||
return out
|
||||
}
|
||||
|
||||
// maxIPv6ExtHeaders caps the extension-header walk in IPv6FindUpperProtocol.
|
||||
// RFC 8200 legal chains are shorter (each header at most once, Destination
|
||||
// Options at most twice), so the cap only bites crafted packets, which would
|
||||
// otherwise make us walk their whole payload 8 bytes at a time.
|
||||
const maxIPv6ExtHeaders = 8
|
||||
|
||||
// IPv6FindUpperProtocol walks packet's IPv6 extension-header chain and
|
||||
// returns the terminal (upper-layer) protocol number, the byte offset where
|
||||
// that protocol's header begins, and whether the packet is a non-first
|
||||
// fragment. It steps over Hop-by-Hop (0), Routing (43), Fragment (44),
|
||||
// AH (51), and Destination Options (60); anything else — including ESP,
|
||||
// whose payload is encrypted — terminates the walk.
|
||||
//
|
||||
// For a non-first fragment, nextHeader still names the flow's upper
|
||||
// protocol (copied from the fragment header) but offset points at fragment
|
||||
// payload, not a real transport header: consult isFragment before
|
||||
// dereferencing offset. If the chain is truncated, over-long, or the packet
|
||||
// is shorter than an IPv6 header, the walk stops early and nextHeader is
|
||||
// whatever it stopped on (59, IPPROTO_NONE, for the too-short case) —
|
||||
// callers treat any non-transport result as unclassifiable.
|
||||
func IPv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool) {
|
||||
if len(packet) < ipv6.HeaderLen {
|
||||
return 59, 0, false // IPPROTO_NONE: nothing to classify
|
||||
}
|
||||
func ipv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool) {
|
||||
nextHeader = packet[6]
|
||||
offset = ipv6.HeaderLen
|
||||
|
||||
for range maxIPv6ExtHeaders {
|
||||
for {
|
||||
switch nextHeader {
|
||||
case 0, 43, 60: // Hop-by-Hop, Routing, Destination
|
||||
if len(packet) < offset+2 {
|
||||
@@ -390,7 +367,6 @@ func IPv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragm
|
||||
return nextHeader, offset, isFragment
|
||||
}
|
||||
}
|
||||
return nextHeader, offset, isFragment
|
||||
}
|
||||
|
||||
func CreateICMPEchoResponse(packet, out []byte) []byte {
|
||||
|
||||
@@ -515,121 +515,3 @@ func TestCreateICMPEchoResponse_IPv6_NotICMPv6(t *testing.T) {
|
||||
result := CreateICMPEchoResponse(packet, out)
|
||||
assert.Nil(t, result)
|
||||
}
|
||||
|
||||
func TestIPv6FindUpperProtocol(t *testing.T) {
|
||||
src := net.ParseIP("fd00::1")
|
||||
dst := net.ParseIP("fd00::2")
|
||||
|
||||
// extHdr builds one 8-byte-unit extension header: next, hdrExtLen
|
||||
// ((extra+1)*8 bytes total), padded to size.
|
||||
extHdr := func(next uint8, extra int) []byte {
|
||||
b := make([]byte, (extra+1)*8)
|
||||
b[0] = next
|
||||
b[1] = uint8(extra)
|
||||
return b
|
||||
}
|
||||
|
||||
t.Run("no extension headers", func(t *testing.T) {
|
||||
for _, proto := range []uint8{6, 17, 58} {
|
||||
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, proto, make([]byte, 20)))
|
||||
assert.Equal(t, proto, nh)
|
||||
assert.Equal(t, ipv6.HeaderLen, offset)
|
||||
assert.False(t, frag)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("hop-by-hop then TCP", func(t *testing.T) {
|
||||
payload := append(extHdr(6, 0), make([]byte, 20)...)
|
||||
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, payload))
|
||||
assert.Equal(t, uint8(6), nh)
|
||||
assert.Equal(t, ipv6.HeaderLen+8, offset)
|
||||
assert.False(t, frag)
|
||||
})
|
||||
|
||||
t.Run("chained headers honor length units", func(t *testing.T) {
|
||||
// Hop-by-Hop (8B) -> Dest Options (16B) -> Routing (8B) -> UDP.
|
||||
payload := extHdr(60, 0)
|
||||
payload = append(payload, extHdr(43, 1)...)
|
||||
payload = append(payload, extHdr(17, 0)...)
|
||||
payload = append(payload, make([]byte, 8)...)
|
||||
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, payload))
|
||||
assert.Equal(t, uint8(17), nh)
|
||||
assert.Equal(t, ipv6.HeaderLen+8+16+8, offset)
|
||||
assert.False(t, frag)
|
||||
})
|
||||
|
||||
t.Run("AH length is in 4-byte units plus 2", func(t *testing.T) {
|
||||
// AH payload-len byte 4 -> (4+2)*4 = 24 bytes on the wire.
|
||||
ah := make([]byte, 24)
|
||||
ah[0] = 6
|
||||
ah[1] = 4
|
||||
payload := append(ah, make([]byte, 20)...)
|
||||
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 51, payload))
|
||||
assert.Equal(t, uint8(6), nh)
|
||||
assert.Equal(t, ipv6.HeaderLen+24, offset)
|
||||
assert.False(t, frag)
|
||||
})
|
||||
|
||||
t.Run("first fragment walks to the transport header", func(t *testing.T) {
|
||||
frag := make([]byte, 8)
|
||||
frag[0] = 17
|
||||
binary.BigEndian.PutUint16(frag[2:4], 0x0001) // offset 0, M=1
|
||||
payload := append(frag, make([]byte, 8)...)
|
||||
nh, offset, isFrag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 44, payload))
|
||||
assert.Equal(t, uint8(17), nh)
|
||||
assert.Equal(t, ipv6.HeaderLen+8, offset)
|
||||
assert.False(t, isFrag, "first fragment carries the real transport header")
|
||||
})
|
||||
|
||||
t.Run("non-first fragment is flagged", func(t *testing.T) {
|
||||
frag := make([]byte, 8)
|
||||
frag[0] = 17
|
||||
binary.BigEndian.PutUint16(frag[2:4], 1<<3) // offset 1, M=0
|
||||
payload := append(frag, make([]byte, 8)...)
|
||||
nh, _, isFrag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 44, payload))
|
||||
assert.Equal(t, uint8(17), nh, "fragment header still names the flow's L4")
|
||||
assert.True(t, isFrag, "offset points at fragment payload, not a header")
|
||||
})
|
||||
|
||||
t.Run("ESP terminates the walk", func(t *testing.T) {
|
||||
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 50, make([]byte, 16)))
|
||||
assert.Equal(t, uint8(50), nh)
|
||||
assert.Equal(t, ipv6.HeaderLen, offset)
|
||||
assert.False(t, frag)
|
||||
})
|
||||
|
||||
t.Run("unknown protocol terminates the walk", func(t *testing.T) {
|
||||
nh, offset, _ := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 132, make([]byte, 16))) // SCTP
|
||||
assert.Equal(t, uint8(132), nh)
|
||||
assert.Equal(t, ipv6.HeaderLen, offset)
|
||||
})
|
||||
|
||||
t.Run("truncated extension header stops the walk", func(t *testing.T) {
|
||||
// Next header says Hop-by-Hop but the packet ends at the IPv6 header.
|
||||
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, nil))
|
||||
assert.Equal(t, uint8(0), nh, "unresolvable chain returns the extension header it stopped on")
|
||||
assert.Equal(t, ipv6.HeaderLen, offset)
|
||||
assert.False(t, frag)
|
||||
})
|
||||
|
||||
t.Run("crafted over-long chain hits the cap", func(t *testing.T) {
|
||||
// Ten chained Hop-by-Hop headers, then TCP. Illegal per RFC 8200
|
||||
// (Hop-by-Hop may only appear first); the cap must stop the walk
|
||||
// before it resolves rather than crawling arbitrary crafted chains.
|
||||
var payload []byte
|
||||
for i := 0; i < 9; i++ {
|
||||
payload = append(payload, extHdr(0, 0)...)
|
||||
}
|
||||
payload = append(payload, extHdr(6, 0)...)
|
||||
payload = append(payload, make([]byte, 20)...)
|
||||
nh, _, _ := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, payload))
|
||||
assert.Equal(t, uint8(0), nh, "walk must stop at the cap, not resolve to TCP")
|
||||
})
|
||||
|
||||
t.Run("packet shorter than an IPv6 header", func(t *testing.T) {
|
||||
nh, offset, frag := IPv6FindUpperProtocol(make([]byte, 39))
|
||||
assert.Equal(t, uint8(59), nh) // IPPROTO_NONE
|
||||
assert.Equal(t, 0, offset)
|
||||
assert.False(t, frag)
|
||||
})
|
||||
}
|
||||
|
||||
+1
-9
@@ -36,10 +36,6 @@ type LightHouse struct {
|
||||
myVpnNetworksTable *bart.Lite
|
||||
punchy *Punchy
|
||||
|
||||
// localAddrsFn enumerates the underlay addresses we advertise. It is a field so tests can supply simulated
|
||||
// addresses rather than whatever this machine's NICs happen to be. Set it before Start.
|
||||
localAddrsFn func(*LocalAllowList) []netip.Addr
|
||||
|
||||
// Local cache of answers from light houses
|
||||
// map of vpn addr to answers
|
||||
addrMap map[netip.Addr]*RemoteList
|
||||
@@ -111,10 +107,6 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
|
||||
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
|
||||
l: l,
|
||||
}
|
||||
h.localAddrsFn = func(al *LocalAllowList) []netip.Addr {
|
||||
return localAddrs(h.l, al)
|
||||
}
|
||||
|
||||
lighthouses := make([]netip.Addr, 0)
|
||||
h.lighthouses.Store(&lighthouses)
|
||||
staticList := make(map[netip.Addr]struct{})
|
||||
@@ -926,7 +918,7 @@ func (lh *LightHouse) SendUpdate() {
|
||||
}
|
||||
|
||||
lal := lh.GetLocalAllowList()
|
||||
for _, e := range lh.localAddrsFn(lal) {
|
||||
for _, e := range localAddrs(lh.l, lal) {
|
||||
if lh.myVpnNetworksTable.Contains(e) {
|
||||
continue
|
||||
}
|
||||
|
||||
+1
-1
@@ -498,7 +498,7 @@ type testEncWriter struct {
|
||||
protocolVersion cert.Version
|
||||
}
|
||||
|
||||
func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
|
||||
func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool) {
|
||||
}
|
||||
func (tw *testEncWriter) Handshake(vpnIp netip.Addr) {
|
||||
}
|
||||
|
||||
@@ -6,14 +6,12 @@ import (
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"runtime/debug"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/cpupick"
|
||||
"github.com/slackhq/nebula/overlay"
|
||||
"github.com/slackhq/nebula/sshd"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
@@ -167,13 +165,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
|
||||
for i := 0; i < routines; i++ {
|
||||
l.Info("listening", "addr", netip.AddrPortFrom(listenHost, uint16(port)))
|
||||
batchSize := c.GetInt("listen.batch", 64)
|
||||
if batchSize < 1 {
|
||||
oldBatch := batchSize
|
||||
batchSize = 1
|
||||
l.Warn("listen.batch size is invalid", "provided", oldBatch, "overridden to", batchSize)
|
||||
}
|
||||
udpServer, err := udp.NewListener(l, listenHost, port, routines > 1, batchSize)
|
||||
udpServer, err := udp.NewListener(l, listenHost, port, routines > 1, c.GetInt("listen.batch", 64))
|
||||
if err != nil {
|
||||
return nil, util.NewContextualError("Failed to open udp listener", m{"queue": i}, err)
|
||||
}
|
||||
@@ -224,17 +216,8 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
|
||||
pinThreads := c.GetBool("tun.pin_threads", true)
|
||||
cpuAffinity := parseCpuAffinity(c, l, routines)
|
||||
if pinThreads && routines > 1 && len(cpuAffinity) == 0 && !configTest {
|
||||
// The operator didn't choose pin CPUs, so pick a default set that
|
||||
// prefers performance cores and doesn't stack co-located instances
|
||||
// onto allowed[0]. The bound UDP port keys the per-instance spread:
|
||||
// distinct across instances sharing a box, stable across restarts.
|
||||
// A nil result keeps listenIn's stock allowed[i] fallback.
|
||||
key := uint64(os.Getpid())
|
||||
if ap, err := udpConns[0].LocalAddr(); err == nil && ap.Port() != 0 {
|
||||
key = uint64(ap.Port())
|
||||
}
|
||||
cpuAffinity = cpupick.Default(routines, key, l)
|
||||
if pinThreads && len(cpuAffinity) == 0 && !configTest {
|
||||
cpuAffinity = defaultCPUAffinityAvoidingIRQs(l, routines)
|
||||
}
|
||||
|
||||
ifConfig := &InterfaceConfig{
|
||||
@@ -277,6 +260,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
ifce.reloadDisconnectInvalid(c)
|
||||
ifce.reloadSendRecvError(c)
|
||||
ifce.reloadAcceptRecvError(c)
|
||||
ifce.reloadEcn(c)
|
||||
|
||||
handshakeManager.f = ifce
|
||||
go handshakeManager.Run(ctx)
|
||||
@@ -297,8 +281,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
|
||||
attachCommands(l, c, ssh, ifce)
|
||||
|
||||
networkChanges := udp.NewNetworkChangeMonitor(ctx, l, c)
|
||||
|
||||
return &Control{
|
||||
state: StateReady,
|
||||
f: ifce,
|
||||
@@ -309,7 +291,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
statsStart: stats.Start,
|
||||
dnsStart: ds.Start,
|
||||
lighthouseStart: lightHouse.StartUpdateWorker,
|
||||
networkChangeStart: networkChanges.Start,
|
||||
connectionManagerStart: connManager.Start,
|
||||
}, nil
|
||||
}
|
||||
@@ -378,6 +359,57 @@ func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int {
|
||||
return cpus
|
||||
}
|
||||
|
||||
// defaultCPUAffinityAvoidingIRQs picks the default pin set for the tun
|
||||
// readers when tun.cpu_affinity is unset: allowed CPUs that do NOT service
|
||||
// any physical NIC's interrupts. The stock allowed[i] spread pins the
|
||||
// encrypt threads onto exactly the cores most drivers affine their first RX
|
||||
// queue IRQs to, so whenever a flow's RSS queue fires on a core hosting a
|
||||
// tun reader, NAPI and encrypt fight for the core and per-flow throughput
|
||||
// drops (measured: REV 8.4 vs 10.2 Gbps on the same hardware, 2026-07-14).
|
||||
//
|
||||
// Returns nil — keeping the old allowed[i] fallback in listenIn — when IRQ
|
||||
// info is unavailable or when there aren't enough IRQ-free CPUs to give
|
||||
// every routine its own core: silently doubling readers up on fewer cores
|
||||
// is worse than the occasional IRQ collision. NICs whose vectors blanket
|
||||
// every CPU (e.g. mlx5 defaults to one queue per core) make avoidance
|
||||
// impossible; narrowing the NIC's spread (ethtool -X <dev> equal N, or
|
||||
// /proc/irq/*/smp_affinity) or setting tun.cpu_affinity explicitly makes it
|
||||
// effective.
|
||||
func defaultCPUAffinityAvoidingIRQs(l *slog.Logger, routines int) []int {
|
||||
irq, err := util.NICIRQCPUs()
|
||||
if err != nil || len(irq) == 0 {
|
||||
return nil
|
||||
}
|
||||
allowed, err := util.AllowedCPUs()
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
cpus := chooseIRQFreeCPUs(allowed, irq, routines)
|
||||
if cpus == nil {
|
||||
l.Info("not enough CPUs are free of NIC IRQs to give every tun reader its own; using the default spread",
|
||||
"routines", routines, "allowed", len(allowed), "irqCPUs", len(irq))
|
||||
return nil
|
||||
}
|
||||
l.Info("pinning tun readers to CPUs clear of NIC IRQs", "cpus", cpus)
|
||||
return cpus
|
||||
}
|
||||
|
||||
// chooseIRQFreeCPUs returns the first `routines` allowed CPUs not present in
|
||||
// irq, or nil if fewer than `routines` qualify.
|
||||
func chooseIRQFreeCPUs(allowed []int, irq map[int]bool, routines int) []int {
|
||||
free := make([]int, 0, routines)
|
||||
for _, cpu := range allowed {
|
||||
if irq[cpu] {
|
||||
continue
|
||||
}
|
||||
free = append(free, cpu)
|
||||
if len(free) == routines {
|
||||
return free
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func moduleVersion() string {
|
||||
info, ok := debug.ReadBuildInfo()
|
||||
if !ok {
|
||||
|
||||
@@ -9,6 +9,26 @@ import (
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestChooseIRQFreeCPUs(t *testing.T) {
|
||||
irq := map[int]bool{0: true, 1: true, 2: true, 3: true}
|
||||
|
||||
// Plenty of IRQ-free CPUs: take the first `routines` of them in order.
|
||||
assert.Equal(t, []int{4, 5}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 2))
|
||||
|
||||
// Exactly enough.
|
||||
assert.Equal(t, []int{4, 5, 6}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 3))
|
||||
|
||||
// Not enough IRQ-free CPUs: nil, caller keeps the old default rather
|
||||
// than doubling readers up on shared cores.
|
||||
assert.Nil(t, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4}, irq, 2))
|
||||
|
||||
// No IRQ info at all behaves like a plain prefix of allowed.
|
||||
assert.Equal(t, []int{0, 1}, chooseIRQFreeCPUs([]int{0, 1, 2}, map[int]bool{}, 2))
|
||||
|
||||
// Non-contiguous allowed set (cgroup cpuset) with holes.
|
||||
assert.Equal(t, []int{9, 12}, chooseIRQFreeCPUs([]int{1, 3, 9, 12}, map[int]bool{1: true, 3: true}, 2))
|
||||
}
|
||||
|
||||
func TestParseCpuAffinity(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
|
||||
|
||||
@@ -164,48 +164,3 @@ func TestCipherStateNilSafety(t *testing.T) {
|
||||
assert.Empty(t, out)
|
||||
assert.Equal(t, 0, cc.Overhead())
|
||||
}
|
||||
|
||||
func TestCipherStateAESGCMInPlaceDecrypt(t *testing.T) {
|
||||
enc, dec := buildCipherStates(t, CipherAESGCM)
|
||||
inPlaceDecrypt(t, NewCipherStateAESGCM(enc), NewCipherStateAESGCM(dec))
|
||||
}
|
||||
|
||||
func TestCipherStateChaChaPolyInPlaceDecrypt(t *testing.T) {
|
||||
enc, dec := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||
inPlaceDecrypt(t, NewCipherStateChaChaPoly(enc), NewCipherStateChaChaPoly(dec))
|
||||
}
|
||||
|
||||
func inPlaceDecrypt(t *testing.T, enc, dec CipherState) {
|
||||
t.Helper()
|
||||
const hdrLen = 16
|
||||
plaintext := []byte("in-place decrypt should replace the ciphertext bytes")
|
||||
nb := make([]byte, 12)
|
||||
|
||||
// packet = [16-byte header | ciphertext+tag], like a nebula Message.
|
||||
packet := make([]byte, hdrLen, hdrLen+len(plaintext)+enc.Overhead())
|
||||
for i := range packet {
|
||||
packet[i] = byte(i)
|
||||
}
|
||||
packet, err := enc.EncryptDanger(packet, packet[:hdrLen], plaintext, 1, nb)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Simulate a GRO row: [packet | next segment]. A failed auth on packet
|
||||
// may zero packet's plaintext region but must not touch the header, the
|
||||
// tag, or the neighboring segment.
|
||||
neighbor := []byte("next coalesced segment, must stay intact")
|
||||
row := append(append([]byte(nil), packet...), neighbor...)
|
||||
tampered := row[:len(packet)]
|
||||
tampered[hdrLen] ^= 0x01
|
||||
_, err = dec.DecryptDanger(tampered[hdrLen:hdrLen], tampered[:hdrLen], tampered[hdrLen:], 1, nb)
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, packet[:hdrLen], tampered[:hdrLen], "failed auth must not touch the header")
|
||||
assert.Equal(t, packet[len(packet)-dec.Overhead():], tampered[len(tampered)-dec.Overhead():],
|
||||
"failed auth must not touch the tag")
|
||||
assert.Equal(t, neighbor, row[len(packet):], "failed auth must not touch the next segment")
|
||||
|
||||
out, err := dec.DecryptDanger(packet[hdrLen:hdrLen], packet[:hdrLen], packet[hdrLen:], 1, nb)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, plaintext, out)
|
||||
// The plaintext must be IN the packet buffer, not a fresh allocation.
|
||||
assert.Equal(t, &packet[hdrLen], &out[0], "plaintext must alias the packet buffer")
|
||||
}
|
||||
|
||||
+155
-76
@@ -13,7 +13,7 @@ import (
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/overlay/batch"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
"golang.org/x/net/ipv4"
|
||||
)
|
||||
|
||||
@@ -23,11 +23,7 @@ const (
|
||||
|
||||
var ErrOutOfWindow = errors.New("out of window packet")
|
||||
|
||||
// readOutsidePackets processes one received underlay packet.
|
||||
// Message payloads are decrypted IN PLACE, so packet must stay untouched
|
||||
// by the caller until the batcher for queue q has been flushed
|
||||
func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxContext) {
|
||||
h := rxc.h
|
||||
func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
|
||||
err := h.Parse(packet)
|
||||
if err != nil {
|
||||
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
||||
@@ -95,7 +91,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxCont
|
||||
if isMessageRelay {
|
||||
hostinfo = f.hostMap.QueryRelayIndex(h.RemoteIndex)
|
||||
} else {
|
||||
hostinfo = f.hostMap.QueryIndexCached(h.RemoteIndex, rxc.hostmapCache)
|
||||
hostinfo = f.hostMap.QueryIndex(h.RemoteIndex)
|
||||
}
|
||||
|
||||
// At this point we should have a valid existing tunnel, verify and send
|
||||
@@ -107,32 +103,26 @@ func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxCont
|
||||
return
|
||||
}
|
||||
|
||||
if len(packet) < header.Len+hostinfo.ConnectionState.dKey.Overhead() {
|
||||
f.messageMetrics.RxInvalid(1)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("packet too small", "from", via, "length", len(packet))
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// All remaining packets are encrypted
|
||||
if isMessageRelay {
|
||||
// Relay packets are special, this branch should always early-return
|
||||
err = hostinfo.ConnectionState.VerifyRelay(f.l, h.MessageCounter, packet, rxc.nb)
|
||||
if err != nil {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("Failed to verify relay packet", "error", err, "from", via, "header", h)
|
||||
}
|
||||
return
|
||||
}
|
||||
f.handleOutsideRelayPacket(hostinfo, via, packet, rxc)
|
||||
ci := hostinfo.ConnectionState
|
||||
if !ci.window.Check(f.l, h.MessageCounter) {
|
||||
return
|
||||
}
|
||||
|
||||
out, err := hostinfo.ConnectionState.Decrypt(f.l, h.MessageCounter, packet, rxc.nb)
|
||||
// Relay packets are special
|
||||
if isMessageRelay {
|
||||
f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, lhf, nb, q, localCache, meta)
|
||||
return
|
||||
}
|
||||
|
||||
out, err = f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
||||
if err != nil {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("Failed to decrypt packet", "error", err, "from", via, "header", h)
|
||||
hostinfo.logger(f.l).Debug("Failed to decrypt packet",
|
||||
"error", err,
|
||||
"from", via,
|
||||
"header", h,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
@@ -145,7 +135,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxCont
|
||||
case header.Message:
|
||||
switch h.Subtype {
|
||||
case header.MessageNone:
|
||||
f.handleOutsideMessagePacket(hostinfo, h.MessageCounter, out, rxc)
|
||||
f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, nb, q, localCache, meta)
|
||||
default:
|
||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
|
||||
return
|
||||
@@ -153,23 +143,15 @@ func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxCont
|
||||
|
||||
case header.LightHouse:
|
||||
//TODO: assert via is not relayed
|
||||
rxc.lhh.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f)
|
||||
lhf.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f)
|
||||
|
||||
case header.Test:
|
||||
switch h.Subtype {
|
||||
case header.TestReply:
|
||||
// No-op, useful for the Roaming and connectionManager side-effects above
|
||||
case header.TestRequest:
|
||||
const maxCipherOverhead = 16 //todo we use this too often, needs a real importable const
|
||||
const maxOverhead = header.Len + header.Len + maxCipherOverhead + maxCipherOverhead
|
||||
if maxOverhead+len(out) > len(rxc.scratch) {
|
||||
// A reply that cannot fit in scratch is dropped no matter the log level.
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("dropping oversized test request", "payloadLen", len(out), "from", via)
|
||||
}
|
||||
return
|
||||
}
|
||||
f.send(header.Test, header.TestReply, hostinfo.ConnectionState, hostinfo, out, rxc.nb, rxc.scratch[:0])
|
||||
//recycle the input packet ciphertext as our output buffer
|
||||
f.send(header.Test, header.TestReply, ci, hostinfo, out, nb, packet)
|
||||
default:
|
||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h)
|
||||
return
|
||||
@@ -187,10 +169,28 @@ func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxCont
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, packet []byte, rxc *rxContext) {
|
||||
h := rxc.h
|
||||
// Successfully validated the thing. Get rid of the Relay header and the AEAD tag
|
||||
signedPayload := packet[header.Len : len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
|
||||
// The entire body is sent as AD, not encrypted.
|
||||
// The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value.
|
||||
// The packet is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's
|
||||
// otherwise malformed (meaning, there is no trailing 16 byte AEAD value), then this will result in at worst a 0-length slice
|
||||
// which will gracefully fail in the DecryptDanger call.
|
||||
signedPayload := packet[:len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||
signatureValue := packet[len(packet)-hostinfo.ConnectionState.dKey.Overhead():]
|
||||
var err error
|
||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, signedPayload, signatureValue, h.MessageCounter, nb)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
// Advance the replay window now that the frame is authenticated
|
||||
if !hostinfo.ConnectionState.window.Update(f.l, h.MessageCounter) {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("dropping out of window relay packet", "header", h)
|
||||
}
|
||||
return
|
||||
}
|
||||
// Successfully validated the thing. Get rid of the Relay header.
|
||||
signedPayload = signedPayload[header.Len:]
|
||||
// Pull the Roaming parts up here, and return in all call paths.
|
||||
f.handleHostRoaming(hostinfo, via)
|
||||
// Track usage of both the HostInfo and the Relay for the received & authenticated packet
|
||||
@@ -201,7 +201,9 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
||||
if !ok {
|
||||
// The only way this happens is if hostmap has an index to the correct HostInfo, but the HostInfo is missing
|
||||
// its internal mapping. This should never happen.
|
||||
hostinfo.logger(f.l).Error("HostInfo missing remote relay index", "relayRemoteIndex", h.RemoteIndex)
|
||||
hostinfo.logger(f.l).Error("HostInfo missing remote relay index",
|
||||
"relayRemoteIndex", h.RemoteIndex,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -212,10 +214,11 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
||||
via = ViaSender{
|
||||
UdpAddr: via.UdpAddr,
|
||||
relayHI: hostinfo,
|
||||
remoteIdx: relay.RemoteIndex,
|
||||
relay: relay,
|
||||
IsRelayed: true,
|
||||
}
|
||||
f.readOutsidePackets(via, signedPayload, rxc)
|
||||
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache, meta)
|
||||
case ForwardingType:
|
||||
// Find the target HostInfo relay object
|
||||
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
||||
@@ -232,11 +235,9 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
||||
if targetRelay.State == Established {
|
||||
switch targetRelay.Type {
|
||||
case ForwardingType:
|
||||
// Forward this packet through the relay tunnel, rebuilding it in place.
|
||||
// Encode overwrites the old outer header, and the new AEAD tag lands where the old one was
|
||||
fwdBuf := packet[:0]
|
||||
//todo it would potentially be nice to batch these
|
||||
f.SendVia(targetHI, targetRelay, signedPayload, rxc.nb, fwdBuf, true, rxc.q)
|
||||
// Forward this packet through the relay tunnel
|
||||
// Find the target HostInfo //todo it would potentially be nice to batch these
|
||||
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
||||
case TerminalType:
|
||||
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
||||
return
|
||||
@@ -317,11 +318,7 @@ var (
|
||||
)
|
||||
|
||||
// newPacket validates and parses the interesting bits for the firewall out of the ip and sub protocol headers
|
||||
func newPacket(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
// fp is reused across packets; reset the parse byproducts so an early-error return cannot
|
||||
// leak the previous packet's offsets.
|
||||
fp.IPHdrLen = 0
|
||||
fp.FragAny = false
|
||||
func newPacket(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
if len(data) < 1 {
|
||||
return ErrPacketTooShort
|
||||
}
|
||||
@@ -336,7 +333,7 @@ func newPacket(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
return ErrUnknownIPVersion
|
||||
}
|
||||
|
||||
func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
dataLen := len(data)
|
||||
if dataLen < ipv6.HeaderLen {
|
||||
return ErrIPv6PacketTooShort
|
||||
@@ -362,7 +359,6 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
switch proto {
|
||||
case layers.IPProtocolESP, layers.IPProtocolNoNextHeader:
|
||||
fp.Protocol = uint8(proto)
|
||||
fp.IPHdrLen = offset
|
||||
fp.RemotePort = 0
|
||||
fp.LocalPort = 0
|
||||
fp.Fragment = false
|
||||
@@ -373,7 +369,6 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
return ErrIPv6PacketTooShort
|
||||
}
|
||||
fp.Protocol = uint8(proto)
|
||||
fp.IPHdrLen = offset
|
||||
fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6
|
||||
icmptype := data[offset+1]
|
||||
switch icmptype {
|
||||
@@ -391,9 +386,6 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
}
|
||||
|
||||
fp.Protocol = uint8(proto)
|
||||
// offset is the L4 header start: 40 for a plain packet, past the extension chain
|
||||
// otherwise. The coalescer only accepts 40.
|
||||
fp.IPHdrLen = offset
|
||||
if incoming {
|
||||
fp.RemotePort = binary.BigEndian.Uint16(data[offset : offset+2])
|
||||
fp.LocalPort = binary.BigEndian.Uint16(data[offset+2 : offset+4])
|
||||
@@ -411,9 +403,6 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
return ErrIPv6PacketTooShort
|
||||
}
|
||||
|
||||
// A fragment shape the coalescer must not touch either way, first fragment included.
|
||||
fp.FragAny = true
|
||||
|
||||
// Check if this is the first fragment
|
||||
fragmentOffset := binary.BigEndian.Uint16(data[offset+2:offset+4]) &^ uint16(0x7) // Remove the reserved and M flag bits
|
||||
if fragmentOffset != 0 {
|
||||
@@ -455,7 +444,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
return ErrIPv6CouldNotFindPayload
|
||||
}
|
||||
|
||||
func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
// Do we at least have an ipv4 header worth of data?
|
||||
if len(data) < ipv4.HeaderLen {
|
||||
return ErrIPv4PacketTooShort
|
||||
@@ -472,10 +461,6 @@ func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
// Check if this is the second or further fragment of a fragmented packet.
|
||||
flagsfrags := binary.BigEndian.Uint16(data[6:8])
|
||||
fp.Fragment = (flagsfrags & 0x1FFF) != 0
|
||||
// Any fragmentation at all (MF or offset): first fragments have readable ports for the
|
||||
// firewall but must never be coalesced.
|
||||
fp.FragAny = (flagsfrags & 0x3fff) != 0
|
||||
fp.IPHdrLen = ihl
|
||||
|
||||
// Firewall handles protocol checks
|
||||
fp.Protocol = data[9]
|
||||
@@ -519,23 +504,117 @@ func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, messageCounter uint64, out []byte, rxc *rxContext) {
|
||||
err := newPacket(out, true, rxc.fwPacket)
|
||||
func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, out []byte, packet []byte, h *header.H, nb []byte) ([]byte, error) {
|
||||
var err error
|
||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], mc, nb)
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).Warn("Error while validating inbound packet", "error", err, "packet", out)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if !hostinfo.ConnectionState.window.Update(f.l, mc) {
|
||||
return nil, ErrOutOfWindow
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// 2-bit IP-level ECN codepoints (lower bits of IPv4 ToS / IPv6 TC).
|
||||
const (
|
||||
ecnNotECT = 0x00
|
||||
ecnECT1 = 0x01
|
||||
ecnECT0 = 0x02
|
||||
ecnCE = 0x03
|
||||
)
|
||||
|
||||
// applyOuterECN folds an outer CE mark from the underlay into the inner
|
||||
// IP header per RFC 6040 normal mode. It mutates pkt[1] in place. Other
|
||||
// codepoints are advisory only and leave the inner unchanged.
|
||||
//
|
||||
// Merge cases (outer × inner → action):
|
||||
//
|
||||
// outer != CE : no-op (inner is authoritative)
|
||||
// outer == CE, inner Not-ECT : log; cannot propagate to a non-ECN host
|
||||
// outer == CE, inner ECT/CE : rewrite inner ECN to CE
|
||||
func applyOuterECN(pkt []byte, outerECN byte, hostinfo *HostInfo, l *slog.Logger) {
|
||||
if outerECN&ecnCE != ecnCE || len(pkt) < 2 {
|
||||
return
|
||||
}
|
||||
switch pkt[0] >> 4 {
|
||||
case 4:
|
||||
switch pkt[1] & 0x03 {
|
||||
case ecnNotECT:
|
||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
|
||||
}
|
||||
case ecnCE:
|
||||
// Already CE.
|
||||
default:
|
||||
// Rewriting the ToS byte invalidates the IPv4 header checksum, so
|
||||
// patch it incrementally per RFC 1624 (HC' = ~(~HC + ~m + m')). The
|
||||
// ToS is the low byte of the 16-bit word at pkt[0:2]; the header
|
||||
// checksum lives at pkt[10:12]. A header too short to carry a
|
||||
// checksum can't be fixed up here, so leave it for newPacket to
|
||||
// reject rather than emit a mangled packet.
|
||||
if len(pkt) < ipv4.HeaderLen {
|
||||
return
|
||||
}
|
||||
m := binary.BigEndian.Uint16(pkt[0:2])
|
||||
pkt[1] = (pkt[1] &^ 0x03) | ecnCE
|
||||
mNew := binary.BigEndian.Uint16(pkt[0:2])
|
||||
sum := uint32(^binary.BigEndian.Uint16(pkt[10:12])) + uint32(^m) + uint32(mNew)
|
||||
for sum > 0xffff {
|
||||
sum = (sum >> 16) + (sum & 0xffff)
|
||||
}
|
||||
binary.BigEndian.PutUint16(pkt[10:12], ^uint16(sum))
|
||||
}
|
||||
case 6:
|
||||
switch (pkt[1] >> 4) & 0x03 {
|
||||
case ecnNotECT:
|
||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
|
||||
}
|
||||
case ecnCE:
|
||||
// Already CE.
|
||||
default:
|
||||
pkt[1] = (pkt[1] &^ 0x30) | (ecnCE << 4)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
|
||||
// RFC 6040 normal-mode combine: fold any outer CE mark stamped by the
|
||||
// underlay into the inner header before firewall + TUN write. Other
|
||||
// outer codepoints are advisory only — we keep the inner unchanged.
|
||||
if f.ecnEnabled.Load() {
|
||||
applyOuterECN(out, meta.OuterECN, hostinfo, f.l)
|
||||
}
|
||||
|
||||
err := newPacket(out, true, fwPacket)
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
|
||||
"error", err,
|
||||
"packet", out,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
dropReason := f.firewall.Drop(rxc.fwPacket.Packet, true, hostinfo, f.pki.GetCAPool(), rxc.ctCache.Get())
|
||||
dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache)
|
||||
if dropReason != nil {
|
||||
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, rxc.nb, rxc.scratch, rxc.q)
|
||||
// NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore
|
||||
// This gives us a buffer to build the reject packet in. With UDP GRO this is a single segment of a shared
|
||||
// recvmmsg row whose capacity runs to the end of the whole row, so cap it to its own length (cap==len) to
|
||||
// keep the reject builder from writing past this segment into the next, not-yet-processed coalesced segment.
|
||||
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet[:len(packet):len(packet)], q)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("dropping inbound packet", "fwPacket", rxc.fwPacket, "reason", dropReason)
|
||||
hostinfo.logger(f.l).Debug("dropping inbound packet",
|
||||
"fwPacket", fwPacket,
|
||||
"reason", dropReason,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
err = f.batchers[rxc.q].Commit(out, batch.SortKey{Epoch: hostinfo.ConnectionState.epoch, Counter: messageCounter}, rxc.fwPacket)
|
||||
err = f.batchers[q].Commit(out)
|
||||
if err != nil {
|
||||
f.l.Error("Failed to write to tun", "error", err)
|
||||
}
|
||||
|
||||
+5
-90
@@ -17,7 +17,7 @@ import (
|
||||
)
|
||||
|
||||
func Test_newPacket(t *testing.T) {
|
||||
p := &firewall.ParsedPacket{}
|
||||
p := &firewall.Packet{}
|
||||
|
||||
// length fails
|
||||
err := newPacket([]byte{}, true, p)
|
||||
@@ -96,7 +96,7 @@ func Test_newPacket(t *testing.T) {
|
||||
}
|
||||
|
||||
func Test_newPacket_v6(t *testing.T) {
|
||||
p := &firewall.ParsedPacket{}
|
||||
p := &firewall.Packet{}
|
||||
|
||||
// invalid ipv6
|
||||
ip := layers.IPv6{
|
||||
@@ -345,7 +345,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
}
|
||||
|
||||
func Test_newPacket_ipv6Fragment(t *testing.T) {
|
||||
p := &firewall.ParsedPacket{}
|
||||
p := &firewall.Packet{}
|
||||
|
||||
ip := &layers.IPv6{
|
||||
Version: 6,
|
||||
@@ -525,7 +525,7 @@ func BenchmarkParseV6(b *testing.B) {
|
||||
secondFrag = append(secondFrag, fragHeader...)
|
||||
secondFrag = append(secondFrag, []byte{0xde, 0xad, 0xbe, 0xef}...)
|
||||
|
||||
fp := &firewall.ParsedPacket{}
|
||||
fp := &firewall.Packet{}
|
||||
|
||||
b.Run("Normal", func(b *testing.B) {
|
||||
for i := 0; i < b.N; i++ {
|
||||
@@ -649,7 +649,7 @@ func serializeAH(ah *layers.IPSecAH) []byte {
|
||||
// host OS parses the real header, a firewall port/proto bypass. The fix makes parseV6 land
|
||||
// on the same offset the host does.
|
||||
func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
|
||||
p := &firewall.ParsedPacket{}
|
||||
p := &firewall.Packet{}
|
||||
|
||||
const (
|
||||
hdrLen = 40 // IPv6 header
|
||||
@@ -675,88 +675,3 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
|
||||
// the host delivers to, not the forged 443 at the overflowed offset.
|
||||
assert.Equal(t, uint16(22), p.LocalPort, "firewall must parse the real transport header, not the overflowed offset")
|
||||
}
|
||||
|
||||
// Test_newPacket_parsedFields pins the ParsedPacket byproducts the RX
|
||||
// batcher consumes: IPHdrLen (the true L4 offset) and FragAny (any fragment
|
||||
// shape at all — unlike Packet.Fragment, which is port-oriented and true
|
||||
// only for non-first fragments).
|
||||
func Test_newPacket_parsedFields(t *testing.T) {
|
||||
p := &firewall.ParsedPacket{}
|
||||
|
||||
// Plain IPv4 TCP, IHL 20: L4 offset 20, no fragment shape.
|
||||
v4 := make([]byte, 28)
|
||||
v4[0] = 0x45
|
||||
v4[9] = firewall.ProtoTCP
|
||||
binary.BigEndian.PutUint16(v4[6:8], 0x4000) // DF only
|
||||
require.NoError(t, newPacket(v4, true, p))
|
||||
assert.Equal(t, 20, p.IPHdrLen)
|
||||
assert.False(t, p.FragAny)
|
||||
assert.False(t, p.Fragment)
|
||||
|
||||
// IPv4 first fragment (MF set, offset 0): the firewall can read ports
|
||||
// (Fragment false) but the coalescer must not touch it (FragAny true).
|
||||
ff := make([]byte, 28)
|
||||
ff[0] = 0x45
|
||||
ff[9] = firewall.ProtoUDP
|
||||
binary.BigEndian.PutUint16(ff[6:8], 0x2000) // MF, offset 0
|
||||
require.NoError(t, newPacket(ff, true, p))
|
||||
assert.False(t, p.Fragment)
|
||||
assert.True(t, p.FragAny)
|
||||
assert.Equal(t, 20, p.IPHdrLen)
|
||||
|
||||
// IPv4 non-first fragment (nonzero offset): both flags set.
|
||||
nf := make([]byte, 28)
|
||||
nf[0] = 0x45
|
||||
nf[9] = firewall.ProtoUDP
|
||||
binary.BigEndian.PutUint16(nf[6:8], 0x00b9)
|
||||
require.NoError(t, newPacket(nf, true, p))
|
||||
assert.True(t, p.Fragment)
|
||||
assert.True(t, p.FragAny)
|
||||
|
||||
// IPv4 with options (IHL 24): IPHdrLen tracks the real L4 offset.
|
||||
opts := make([]byte, 32)
|
||||
opts[0] = 0x46
|
||||
opts[9] = firewall.ProtoTCP
|
||||
binary.BigEndian.PutUint16(opts[6:8], 0x4000)
|
||||
require.NoError(t, newPacket(opts, true, p))
|
||||
assert.Equal(t, 24, p.IPHdrLen)
|
||||
assert.False(t, p.FragAny)
|
||||
|
||||
// Plain IPv6 TCP: L4 at 40.
|
||||
v6 := make([]byte, 60)
|
||||
v6[0] = 0x60
|
||||
v6[6] = firewall.ProtoTCP
|
||||
require.NoError(t, newPacket(v6, true, p))
|
||||
assert.Equal(t, 40, p.IPHdrLen)
|
||||
assert.False(t, p.FragAny)
|
||||
|
||||
// IPv6 hop-by-hop then TCP: IPHdrLen lands past the extension header.
|
||||
hbh := make([]byte, 60)
|
||||
hbh[0] = 0x60
|
||||
hbh[6] = 0 // hop-by-hop
|
||||
hbh[40] = firewall.ProtoTCP
|
||||
hbh[41] = 0 // HdrExtLen 0 -> 8-byte header
|
||||
require.NoError(t, newPacket(hbh, true, p))
|
||||
assert.Equal(t, 48, p.IPHdrLen)
|
||||
assert.False(t, p.FragAny)
|
||||
|
||||
// IPv6 first fragment: terminal proto resolved, FragAny set, Fragment not.
|
||||
f6 := make([]byte, 60)
|
||||
f6[0] = 0x60
|
||||
f6[6] = 44 // fragment extension header
|
||||
f6[40] = firewall.ProtoUDP
|
||||
require.NoError(t, newPacket(f6, true, p))
|
||||
assert.True(t, p.FragAny)
|
||||
assert.False(t, p.Fragment)
|
||||
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
||||
|
||||
// IPv6 non-first fragment: both set, walk stops at the fragment header.
|
||||
f6n := make([]byte, 60)
|
||||
f6n[0] = 0x60
|
||||
f6n[6] = 44
|
||||
f6n[40] = firewall.ProtoUDP
|
||||
binary.BigEndian.PutUint16(f6n[42:44], 0x0008)
|
||||
require.NoError(t, newPacket(f6n, true, p))
|
||||
assert.True(t, p.Fragment)
|
||||
assert.True(t, p.FragAny)
|
||||
}
|
||||
|
||||
+25
-8
@@ -1,11 +1,28 @@
|
||||
package batch
|
||||
|
||||
// SortKey identifies a packet's position in its sender's transmission order.
|
||||
// Epoch is a receiver-local ordinal for the tunnel (ConnectionState) that decrypted the packet:
|
||||
// a re-handshake replaces the tunnel outright and the replacement's epoch is higher,
|
||||
// so the old tunnel's packets sort first during the cutover overlap.
|
||||
// Counter is the packet's AEAD message counter within that tunnel.
|
||||
type SortKey struct {
|
||||
Epoch uint64
|
||||
Counter uint64
|
||||
import "net/netip"
|
||||
|
||||
type RxBatcher interface {
|
||||
// Reserve creates a pkt to borrow
|
||||
Reserve(sz int) []byte
|
||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush
|
||||
Commit(pkt []byte) error
|
||||
// Flush emits every queued packet in arrival order.
|
||||
// Returns the first error observed; keeps draining so one bad packet doesn't hold up the rest.
|
||||
// After Flush returns, borrowed payload slices may be recycled.
|
||||
Flush() error
|
||||
}
|
||||
|
||||
type TxBatcher interface {
|
||||
// Reserve creates a pkt to borrow
|
||||
Reserve(sz int) []byte
|
||||
// Commit borrows pkt and records its destination plus the 2-bit
|
||||
// IP-level ECN codepoint to set on the outer (carrier) header. The
|
||||
// caller must keep pkt valid until the next Flush. Pass 0 (Not-ECT)
|
||||
// to leave the outer ECN field unset.
|
||||
Commit(pkt []byte, dst netip.AddrPort, outerECN byte)
|
||||
// Flush emits every queued packet via the underlying batch writer in arrival order.
|
||||
// Returns an errors.Join of one or more errors.
|
||||
// After Flush returns, borrowed payload slices may be recycled.
|
||||
Flush() error
|
||||
}
|
||||
|
||||
@@ -1,187 +0,0 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"math/rand"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// The checksum-seeding helpers feed the virtio NEEDS_CSUM contract: the L4
|
||||
// checksum field is pre-loaded with the folded (not inverted) pseudo-header
|
||||
// sum, and the kernel later adds the L4 byte sum and inverts. A wrong seed
|
||||
// produces packets every receiver silently drops, with nothing failing on
|
||||
// our side — so these tests check the helpers against an independent
|
||||
// RFC 1071 reference built from explicit pseudo-header bytes, never against
|
||||
// the production checksum code.
|
||||
|
||||
// refSum accumulates big-endian 16-bit words of b (odd tail zero-padded)
|
||||
// into a wide one's-complement accumulator.
|
||||
func refSum(b []byte) uint64 {
|
||||
var s uint64
|
||||
for i := 0; i+1 < len(b); i += 2 {
|
||||
s += uint64(b[i])<<8 | uint64(b[i+1])
|
||||
}
|
||||
if len(b)%2 == 1 {
|
||||
s += uint64(b[len(b)-1]) << 8
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// refFold folds a wide one's-complement accumulator to 16 bits.
|
||||
func refFold(s uint64) uint16 {
|
||||
for s>>16 != 0 {
|
||||
s = s&0xffff + s>>16
|
||||
}
|
||||
return uint16(s)
|
||||
}
|
||||
|
||||
func TestFoldOnceNoInvertEdgeCases(t *testing.T) {
|
||||
cases := []uint32{
|
||||
0, 1, 0xffff,
|
||||
0x10000, // single carry
|
||||
0x1fffe, // 0xffff + 0xffff: carry produces another 0xffff
|
||||
0xffff0000, // high half only
|
||||
0xfffeffff, // fold yields 0x1fffd: needs a second fold
|
||||
0xffffffff, // worst case
|
||||
0x00010001, // simple two-word
|
||||
}
|
||||
for _, c := range cases {
|
||||
want := refFold(uint64(c))
|
||||
if got := foldOnceNoInvert(c); got != want {
|
||||
t.Errorf("foldOnceNoInvert(%#x) = %#x, want %#x", c, got, want)
|
||||
}
|
||||
// Folding a folded value must be a no-op.
|
||||
if got := foldOnceNoInvert(uint32(foldOnceNoInvert(c))); got != foldOnceNoInvert(c) {
|
||||
t.Errorf("foldOnceNoInvert not idempotent at %#x", c)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPseudoSumIPv4MatchesReference(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
src, dst [4]byte
|
||||
proto byte
|
||||
l4Len int
|
||||
}{
|
||||
{"simple", [4]byte{10, 0, 0, 1}, [4]byte{10, 0, 0, 2}, 6, 20},
|
||||
{"zero-len", [4]byte{192, 168, 1, 1}, [4]byte{192, 168, 1, 2}, 17, 0},
|
||||
{"max-len", [4]byte{1, 2, 3, 4}, [4]byte{5, 6, 7, 8}, 6, 65535},
|
||||
{"carry-heavy", [4]byte{255, 255, 255, 255}, [4]byte{255, 255, 255, 254}, 17, 65535},
|
||||
{"broadcastish", [4]byte{255, 255, 255, 255}, [4]byte{255, 255, 255, 255}, 255, 65535},
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
// RFC 793 pseudo-header: src(4) dst(4) zero(1) proto(1) len(2).
|
||||
ph := make([]byte, 12)
|
||||
copy(ph[0:4], c.src[:])
|
||||
copy(ph[4:8], c.dst[:])
|
||||
ph[9] = c.proto
|
||||
binary.BigEndian.PutUint16(ph[10:12], uint16(c.l4Len))
|
||||
want := refFold(refSum(ph))
|
||||
|
||||
got := foldOnceNoInvert(pseudoSumIPv4(c.src[:], c.dst[:], c.proto, c.l4Len))
|
||||
if got != want {
|
||||
t.Errorf("fold(pseudoSumIPv4) = %#x, want %#x", got, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPseudoSumIPv6MatchesReference(t *testing.T) {
|
||||
ones := func(b byte) (a [16]byte) {
|
||||
for i := range a {
|
||||
a[i] = b
|
||||
}
|
||||
return
|
||||
}
|
||||
cases := []struct {
|
||||
name string
|
||||
src, dst [16]byte
|
||||
proto byte
|
||||
l4Len int
|
||||
}{
|
||||
{"simple", [16]byte{0xfe, 0x80, 15: 1}, [16]byte{0xfe, 0x80, 15: 2}, 6, 20},
|
||||
{"zero-len", [16]byte{0x20, 0x01, 15: 9}, [16]byte{0x20, 0x01, 15: 8}, 17, 0},
|
||||
{"max-u16-len", ones(0xff), ones(0xfe), 6, 65535},
|
||||
{"len-past-u16", ones(0xff), ones(0xff), 17, 0x12345}, // exercises the 32-bit split
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
// RFC 8200 pseudo-header: src(16) dst(16) len(4) zero(3) next(1).
|
||||
ph := make([]byte, 40)
|
||||
copy(ph[0:16], c.src[:])
|
||||
copy(ph[16:32], c.dst[:])
|
||||
binary.BigEndian.PutUint32(ph[32:36], uint32(c.l4Len))
|
||||
ph[39] = c.proto
|
||||
want := refFold(refSum(ph))
|
||||
|
||||
got := foldOnceNoInvert(pseudoSumIPv6(c.src[:], c.dst[:], c.proto, c.l4Len))
|
||||
if got != want {
|
||||
t.Errorf("fold(pseudoSumIPv6) = %#x, want %#x", got, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIPv4HdrChecksumMatchesReference(t *testing.T) {
|
||||
rng := rand.New(rand.NewSource(0x1791))
|
||||
for _, hdrLen := range []int{20, 24, 40, 60} {
|
||||
for trial := 0; trial < 200; trial++ {
|
||||
hdr := make([]byte, hdrLen)
|
||||
rng.Read(hdr)
|
||||
hdr[0] = 0x40 | byte(hdrLen/4)
|
||||
hdr[10], hdr[11] = 0, 0 // checksum field zeroed, as the contract requires
|
||||
|
||||
want := ^refFold(refSum(hdr))
|
||||
got := ipv4HdrChecksum(hdr)
|
||||
if got != want {
|
||||
t.Fatalf("ipv4HdrChecksum(len=%d trial=%d) = %#x, want %#x", hdrLen, trial, got, want)
|
||||
}
|
||||
|
||||
// Receiver-side property: with the checksum stored, the full
|
||||
// header must sum to all-ones.
|
||||
binary.BigEndian.PutUint16(hdr[10:12], got)
|
||||
if v := refFold(refSum(hdr)); v != 0xffff {
|
||||
t.Fatalf("stored checksum does not validate: full-header fold = %#x", v)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestChecksumSeedReceiverAcceptance is the end-to-end property the helpers
|
||||
// exist for: seed the TCP checksum field with fold(pseudoSum), do what the
|
||||
// kernel's NEEDS_CSUM completion does (one's-complement sum over the L4
|
||||
// bytes including the seed, then invert, then store), and verify the result
|
||||
// the way a receiver does (pseudo-header + L4 must sum to all-ones).
|
||||
func TestChecksumSeedReceiverAcceptance(t *testing.T) {
|
||||
rng := rand.New(rand.NewSource(0x1826))
|
||||
for trial := 0; trial < 200; trial++ {
|
||||
src := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
|
||||
dst := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
|
||||
payLen := rng.Intn(1500)
|
||||
l4 := make([]byte, 20+payLen)
|
||||
rng.Read(l4)
|
||||
|
||||
// Seed exactly as flushSlot does.
|
||||
seed := foldOnceNoInvert(pseudoSumIPv4(src[:], dst[:], 6, len(l4)))
|
||||
binary.BigEndian.PutUint16(l4[16:18], seed)
|
||||
|
||||
// Kernel NEEDS_CSUM completion: sum the L4 region (seed included,
|
||||
// which is equivalent to summing with the field zeroed and folding
|
||||
// the seed in), invert, store.
|
||||
final := ^refFold(refSum(l4[:16]) + uint64(seed) + refSum(l4[18:]))
|
||||
binary.BigEndian.PutUint16(l4[16:18], final)
|
||||
|
||||
// Receiver validation.
|
||||
ph := make([]byte, 12)
|
||||
copy(ph[0:4], src[:])
|
||||
copy(ph[4:8], dst[:])
|
||||
ph[9] = 6
|
||||
binary.BigEndian.PutUint16(ph[10:12], uint16(len(l4)))
|
||||
if v := refFold(refSum(ph) + refSum(l4)); v != 0xffff {
|
||||
t.Fatalf("trial %d: receiver rejects packet: fold = %#x (seed=%#x final=%#x payLen=%d)",
|
||||
trial, v, seed, final, payLen)
|
||||
}
|
||||
}
|
||||
}
|
||||
+122
-102
@@ -8,125 +8,135 @@ import (
|
||||
// flowKey identifies a transport flow by {src, dst, sport, dport, family}.
|
||||
// Comparable, so map lookups and linear scans over the slot list stay tight.
|
||||
// Shared by the TCP and UDP coalescers; each coalescer keeps its own
|
||||
// openSlots map, so a TCP and UDP flow on the same 5-tuple-without-proto never alias.
|
||||
// openSlots map, so a TCP and UDP flow on the same 5-tuple-without-proto
|
||||
// never alias.
|
||||
type flowKey struct {
|
||||
src, dst [16]byte
|
||||
sport, dport uint16
|
||||
isV6 bool
|
||||
}
|
||||
|
||||
// initialSlots is the starting capacity of the slot pool.
|
||||
// One flow per packet is the worst case,
|
||||
// so this matches a typical carrier-side recvmmsg batch on the UDP socket.
|
||||
// initialSlots is the starting capacity of the slot pool. One flow per
|
||||
// packet is the worst case so this matches a typical carrier-side
|
||||
// recvmmsg batch on the encrypted UDP socket.
|
||||
const initialSlots = 64
|
||||
|
||||
// parseIPAt validates the IP header for lane parsing. newPacket already resolved the L4 protocol
|
||||
// and offset for the firewall, so there is no proto sniff here; the caller's ipHdrLen is
|
||||
// cross-checked instead. A plain header (v4 IHL 20, v6 exactly 40) is the only coalesceable
|
||||
// shape. The v6 check is load-bearing: it rejects extension-header packets whose L4 is not at
|
||||
// byte 40.
|
||||
// parsedIP is the IP-level result of parseIPPrologue. The caller layers
|
||||
// L4-specific parsing (TCP / UDP) on top.
|
||||
type parsedIP struct {
|
||||
fk flowKey
|
||||
ipHdrLen int
|
||||
// pkt is the original buffer trimmed to the IP-declared total length.
|
||||
// Anything below the IP layer (transport parsers) should slice into
|
||||
// pkt rather than the unbounded original.
|
||||
pkt []byte
|
||||
}
|
||||
|
||||
// parseIPPrologue extracts the IP-level fields the coalescers care about:
|
||||
// IHL/payload length, version, src/dst addresses, and the L4 protocol byte.
|
||||
// Returns ok=false for malformed input, IPv4 with options or fragmentation,
|
||||
// or IPv6 with extension headers (all rejected by both coalescers in
|
||||
// identical ways before this refactor).
|
||||
//
|
||||
// The prologues fill fk's addresses and family in place (ports belong to the L4 parser; fk must
|
||||
// be zero on entry so the v4 path leaves src[4:]/dst[4:] clear for map equality) and return pkt
|
||||
// trimmed to the IP-declared length. The receiver-as-out-pointer shape is deliberate: these
|
||||
// functions are too big to inline, and returning structs by value put five 64-byte copies on the
|
||||
// per-packet path.
|
||||
func (fk *flowKey) parseIPAt(pkt []byte, ipHdrLen int) ([]byte, bool) {
|
||||
// On success, p.pkt is len-trimmed to the IP-declared length so callers
|
||||
// don't have to repeat the trim. wantProto is the IANA protocol number to
|
||||
// require (6 for TCP, 17 for UDP); ok=false for any other value.
|
||||
func parseIPPrologue(pkt []byte, wantProto byte) (parsedIP, bool) {
|
||||
var p parsedIP
|
||||
if len(pkt) < 20 {
|
||||
return nil, false
|
||||
return p, false
|
||||
}
|
||||
switch pkt[0] >> 4 {
|
||||
v := pkt[0] >> 4
|
||||
switch v {
|
||||
case 4:
|
||||
if ipHdrLen != 20 {
|
||||
return nil, false
|
||||
ihl := int(pkt[0]&0x0f) * 4
|
||||
if ihl != 20 {
|
||||
return p, false
|
||||
}
|
||||
return fk.parseIPv4Prologue(pkt)
|
||||
if pkt[9] != wantProto {
|
||||
return p, false
|
||||
}
|
||||
// Reject actual fragmentation (MF or non-zero frag offset).
|
||||
if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 {
|
||||
return p, false
|
||||
}
|
||||
totalLen := int(binary.BigEndian.Uint16(pkt[2:4]))
|
||||
if totalLen > len(pkt) || totalLen < ihl {
|
||||
return p, false
|
||||
}
|
||||
p.ipHdrLen = 20
|
||||
p.fk.isV6 = false
|
||||
copy(p.fk.src[:4], pkt[12:16])
|
||||
copy(p.fk.dst[:4], pkt[16:20])
|
||||
p.pkt = pkt[:totalLen]
|
||||
case 6:
|
||||
if ipHdrLen != 40 || len(pkt) < 40 {
|
||||
return nil, false
|
||||
if len(pkt) < 40 {
|
||||
return p, false
|
||||
}
|
||||
return fk.parseIPv6Prologue(pkt)
|
||||
if pkt[6] != wantProto {
|
||||
return p, false
|
||||
}
|
||||
payloadLen := int(binary.BigEndian.Uint16(pkt[4:6]))
|
||||
if 40+payloadLen > len(pkt) {
|
||||
return p, false
|
||||
}
|
||||
p.ipHdrLen = 40
|
||||
p.fk.isV6 = true
|
||||
copy(p.fk.src[:], pkt[8:24])
|
||||
copy(p.fk.dst[:], pkt[24:40])
|
||||
p.pkt = pkt[:40+payloadLen]
|
||||
default:
|
||||
return p, false
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// parseIPv4Prologue is the shared IPv4 tail of the prologue entries; the caller has verified
|
||||
// len(pkt) >= 20 and the version.
|
||||
func (fk *flowKey) parseIPv4Prologue(pkt []byte) ([]byte, bool) {
|
||||
ihl := int(pkt[0]&0x0f) * 4
|
||||
if ihl != 20 {
|
||||
return nil, false
|
||||
}
|
||||
// Reject any fragmentation (MF or nonzero offset). The dispatcher already gated FragAny; kept
|
||||
// as defense in depth, since a fragment folded into a superpacket would corrupt reassembly.
|
||||
if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 {
|
||||
return nil, false
|
||||
}
|
||||
totalLen := int(binary.BigEndian.Uint16(pkt[2:4]))
|
||||
if totalLen > len(pkt) || totalLen < ihl {
|
||||
return nil, false
|
||||
}
|
||||
fk.isV6 = false
|
||||
copy(fk.src[:4], pkt[12:16])
|
||||
copy(fk.dst[:4], pkt[16:20])
|
||||
return pkt[:totalLen], true
|
||||
}
|
||||
|
||||
// parseIPv6Prologue is the shared IPv6 tail; the caller has verified len(pkt) >= 40, the version,
|
||||
// and that the L4 header sits at byte 40.
|
||||
func (fk *flowKey) parseIPv6Prologue(pkt []byte) ([]byte, bool) {
|
||||
payloadLen := int(binary.BigEndian.Uint16(pkt[4:6]))
|
||||
if 40+payloadLen > len(pkt) {
|
||||
return nil, false
|
||||
}
|
||||
fk.isV6 = true
|
||||
copy(fk.src[:], pkt[8:24])
|
||||
copy(fk.dst[:], pkt[24:40])
|
||||
return pkt[:40+payloadLen], true
|
||||
return p, true
|
||||
}
|
||||
|
||||
// ipHeadersMatch compares the IP portion of two packet header prefixes for
|
||||
// byte-for-byte equality on every field that must be identical across coalesced segments.
|
||||
// Size/IPID/IPCsum are masked out.
|
||||
// The full DSCP/ECN byte (IPv4 ToS / IPv6 traffic class) is compared, matching Linux kernel GRO:
|
||||
// segments with differing ECN codepoints must not coalesce,
|
||||
// otherwise ORing e.g. ECT(0) with ECT(1) would fabricate a false CE (congestion) mark or mark a Not-ECT flow as ECN-capable.
|
||||
// byte-for-byte equality on every field that must be identical across
|
||||
// coalesced segments. Size/IPID/IPCsum are masked out. The full DSCP/ECN
|
||||
// byte (IPv4 ToS / IPv6 traffic class) is compared, matching Linux kernel
|
||||
// GRO: segments with differing ECN codepoints must not coalesce, otherwise
|
||||
// ORing e.g. ECT(0) with ECT(1) would fabricate a false CE (congestion)
|
||||
// mark or mark a Not-ECT flow as ECN-capable.
|
||||
//
|
||||
// The transport (L4) portion of the header is checked separately by the per-protocol matcher.
|
||||
// The transport (L4) portion of the header is checked separately by the
|
||||
// per-protocol matcher.
|
||||
func ipHeadersMatch(a, b []byte, isV6 bool) bool {
|
||||
if isV6 {
|
||||
// IPv6: [0:4] = version/TC/flow label (TC[1:0] is ECN, so the full TC byte must match),
|
||||
// [6:40] = next_hdr/hop + src + dst. Skip [4:6] payload_len.
|
||||
return bytes.Equal(a[:4], b[:4]) && bytes.Equal(a[6:40], b[6:40])
|
||||
}
|
||||
// IPv4: [0:2] = version/IHL + DSCP|ECN (full ECN byte must match),
|
||||
// [6:10] = flags/fragoff/TTL/proto, [12:20] = src+dst.
|
||||
// Skip [2:4] total len, [4:6] id, [10:12] csum.
|
||||
return bytes.Equal(a[:2], b[:2]) && bytes.Equal(a[6:10], b[6:10]) && bytes.Equal(a[12:20], b[12:20])
|
||||
}
|
||||
|
||||
// ipv4FlagDF is the Don't Fragment bit in the IPv4 flags byte (header byte 6).
|
||||
const ipv4FlagDF = 0x40
|
||||
|
||||
// ipv4CanCoalesceID reports whether an IPv4 packet whose header starts at
|
||||
// nextHdr may join a chain whose seed header is seedHdr as segment index seg
|
||||
// (the seed is segment 0). Kernel GSO re-stamps outgoing segment IDs as
|
||||
// seed_id+n, so coalescing is only transparent when that re-stamp is either
|
||||
// harmless (DF set: RFC 6864 atomic datagrams, the ID carries no meaning) or
|
||||
// reproduces the original IDs exactly (DF clear + IDs already sequential —
|
||||
// the same admission rule kernel GRO applies). Without this, a DF=0 sender
|
||||
// with non-sequential IDs (e.g. OpenBSD's randomized IDs) could have IDs
|
||||
// rewritten into ranges that collide across superpackets, corrupting
|
||||
// reassembly if the packets are fragmented after the TUN write.
|
||||
//
|
||||
// DF itself is guaranteed uniform across a chain by ipHeadersMatch (byte 6
|
||||
// is inside its compared range), so checking the seed's copy suffices.
|
||||
func ipv4CanCoalesceID(seedHdr, nextHdr []byte, seg int) bool {
|
||||
if seedHdr[6]&ipv4FlagDF != 0 {
|
||||
// IPv6: byte 0 = version/TC[7:4], byte 1 = TC[3:0]/flow[19:16],
|
||||
// bytes [2:4] = flow[15:0], [6:8] = next_hdr/hop, [8:40] = src+dst.
|
||||
// Compare byte 1 fully so ECN (TC[1:0]) must match. Skip [4:6] payload_len.
|
||||
if a[0] != b[0] {
|
||||
return false
|
||||
}
|
||||
if a[1] != b[1] {
|
||||
return false
|
||||
}
|
||||
if !bytes.Equal(a[2:4], b[2:4]) {
|
||||
return false
|
||||
}
|
||||
if !bytes.Equal(a[6:40], b[6:40]) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
expect := binary.BigEndian.Uint16(seedHdr[4:6]) + uint16(seg)
|
||||
return binary.BigEndian.Uint16(nextHdr[4:6]) == expect
|
||||
// IPv4: byte 0 = version/IHL, byte 1 = DSCP(6)|ECN(2),
|
||||
// [6:10] flags/fragoff/TTL/proto, [12:20] src+dst.
|
||||
// Compare byte 1 fully so ECN must match.
|
||||
// Skip [2:4] total len, [4:6] id, [10:12] csum.
|
||||
if a[0] != b[0] {
|
||||
return false
|
||||
}
|
||||
if a[1] != b[1] {
|
||||
return false
|
||||
}
|
||||
if !bytes.Equal(a[6:10], b[6:10]) {
|
||||
return false
|
||||
}
|
||||
if !bytes.Equal(a[12:20], b[12:20]) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// Arena is an injectable byte-slab that hands out non-overlapping borrowed
|
||||
@@ -135,14 +145,17 @@ type Arena struct {
|
||||
buf []byte
|
||||
}
|
||||
|
||||
// NewArena returns an Arena with a pre-allocated backing of the given capacity.
|
||||
// NewArena returns an Arena with a pre-allocated backing of the given
|
||||
// capacity. Pass 0 if you don't intend to call Reserve (e.g. a test that
|
||||
// only feeds the coalescer pre-made []byte packets via Commit).
|
||||
func NewArena(capacity int) *Arena {
|
||||
return &Arena{buf: make([]byte, 0, capacity)}
|
||||
}
|
||||
|
||||
// Reserve hands out a non-overlapping sz-byte slice from the arena.
|
||||
// If the request doesn't fit the current backing, a fresh, larger backing is allocated.
|
||||
// Already-borrowed slices reference the old backing and remain valid until Reset.
|
||||
// Reserve hands out a non-overlapping sz-byte slice from the arena. If the
|
||||
// request doesn't fit the current backing, a fresh, larger backing is
|
||||
// allocated; already-borrowed slices reference the old backing and remain
|
||||
// valid until Reset.
|
||||
func (a *Arena) Reserve(sz int) []byte {
|
||||
if len(a.buf)+sz > cap(a.buf) {
|
||||
newCap := max(cap(a.buf)*2, sz)
|
||||
@@ -153,9 +166,16 @@ func (a *Arena) Reserve(sz int) []byte {
|
||||
return a.buf[start : start+sz : start+sz]
|
||||
}
|
||||
|
||||
// Reset releases every slice handed out since the last Reset.
|
||||
// Callers must not use any previously-borrowed slice after this returns.
|
||||
// The underlying backing array is retained so subsequent Reserves don't re-allocate.
|
||||
// Reset releases every slice handed out since the last Reset. Callers must
|
||||
// not use any previously-borrowed slice after this returns. The underlying
|
||||
// backing array is retained so subsequent Reserves don't re-allocate.
|
||||
func (a *Arena) Reset() {
|
||||
a.buf = a.buf[:0]
|
||||
}
|
||||
|
||||
// Reserver hands out an sz-byte slice valid until its Resetter runs.
|
||||
type Reserver func(sz int) []byte
|
||||
|
||||
// Resetter clears all reservations held by a Reserver. Only the arena's
|
||||
// owner holds one; lanes inside a MultiCoalescer get nil.
|
||||
type Resetter func()
|
||||
|
||||
@@ -1,112 +0,0 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/test"
|
||||
)
|
||||
|
||||
// stagePackets builds the stagedPacket entries Commit would have produced, so dispatch benchmarks
|
||||
// bypass staging and the sort entirely.
|
||||
func stagePackets(pkts [][]byte) []stagedPacket {
|
||||
staged := make([]stagedPacket, len(pkts))
|
||||
for i, p := range pkts {
|
||||
pp := testPP(p)
|
||||
staged[i] = stagedPacket{
|
||||
pkt: p,
|
||||
key: SortKey{Epoch: 1, Counter: uint64(i + 1)},
|
||||
proto: pp.Protocol,
|
||||
fragAny: pp.FragAny,
|
||||
ipHdrLen: uint16(pp.IPHdrLen),
|
||||
}
|
||||
}
|
||||
return staged
|
||||
}
|
||||
|
||||
func flushLanes(b *testing.B, m *MultiCoalescer) {
|
||||
b.Helper()
|
||||
if m.tcp != nil {
|
||||
if err := m.tcp.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
if m.udp != nil {
|
||||
if err := m.udp.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := m.pt.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
// runDispatchBench measures dispatch plus the per-batch lane flush: the post-sort half of the
|
||||
// batcher, which is where the production profile concentrates.
|
||||
func runDispatchBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||
b.Helper()
|
||||
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger())
|
||||
staged := stagePackets(pkts)
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkts[0])))
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
if err := m.dispatch(staged[i%len(staged)]); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if (i+1)%batchSize == 0 {
|
||||
flushLanes(b, m)
|
||||
}
|
||||
}
|
||||
b.StopTimer()
|
||||
flushLanes(b, m)
|
||||
}
|
||||
|
||||
// BenchmarkDispatchSingleFlow is the bulk steady state: every packet past the seed appends.
|
||||
func BenchmarkDispatchSingleFlow(b *testing.B) {
|
||||
runDispatchBench(b, buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200), tcpCoalesceMaxSegs)
|
||||
}
|
||||
|
||||
// BenchmarkDispatchInterleaved16 stresses the openSlots map: 16 flows round-robined defeats the
|
||||
// lastSlot cache on every packet.
|
||||
func BenchmarkDispatchInterleaved16(b *testing.B) {
|
||||
pkts := buildTCPv4Interleaved(16, tcpCoalesceMaxSegs, 1200)
|
||||
runDispatchBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkDispatchAckHeavy alternates MSS data with pure ACKs on one flow — the RX shape of a
|
||||
// bidirectional transfer (the peer's data and its ACKs of our data share the tunnel direction).
|
||||
func BenchmarkDispatchAckHeavy(b *testing.B) {
|
||||
pay := make([]byte, 1200)
|
||||
var pkts [][]byte
|
||||
seq := uint32(1000)
|
||||
for range tcpCoalesceMaxSegs / 2 {
|
||||
pkts = append(pkts, buildTCPv4(seq, tcpAck, pay))
|
||||
seq += uint32(len(pay))
|
||||
pkts = append(pkts, buildTCPv4(seq, tcpAck, nil))
|
||||
}
|
||||
runDispatchBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkDispatchUDPFlow is the QUIC-ish bulk UDP shape.
|
||||
func BenchmarkDispatchUDPFlow(b *testing.B) {
|
||||
pay := make([]byte, 1200)
|
||||
pkts := make([][]byte, udpCoalesceMaxSegs)
|
||||
for i := range pkts {
|
||||
pkts[i] = buildUDPv4(2000, 443, pay)
|
||||
}
|
||||
runDispatchBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkDispatchSeedHeavy sets PSH on every packet so each one seeds and immediately closes
|
||||
// its own slot — the small-write RPC shape, and the upper bound on what the seed path (including
|
||||
// the parsedTCP-to-slot field transfer) can cost.
|
||||
func BenchmarkDispatchSeedHeavy(b *testing.B) {
|
||||
pay := make([]byte, 1200)
|
||||
pkts := make([][]byte, tcpCoalesceMaxSegs)
|
||||
seq := uint32(1000)
|
||||
for i := range pkts {
|
||||
pkts[i] = buildTCPv4(seq, tcpAckPsh, pay)
|
||||
seq += uint32(len(pay))
|
||||
}
|
||||
runDispatchBench(b, pkts, len(pkts))
|
||||
}
|
||||
@@ -1,76 +0,0 @@
|
||||
package batch
|
||||
|
||||
//TODO refactor this away
|
||||
// This file holds the lanes' self-parsing Commit entries and the proto-checking parsers behind
|
||||
// them. Production traffic enters the lanes only through MultiCoalescer.dispatch and the At
|
||||
// parsers; these wrappers reproduce that path (including seal-all on unparseable shapes) on top
|
||||
// of a local parse, so tests and benches can drive one lane with nothing but a packet.
|
||||
|
||||
// parseIPPrologue resolves the IP version, requires the L4 protocol to match wantProto (6 TCP,
|
||||
// 17 UDP), and defers to the shared per-version cores. Returns the trimmed packet and the L4
|
||||
// offset; fk must be zero on entry and is filled in place.
|
||||
func (fk *flowKey) parseIPPrologue(pkt []byte, wantProto byte) ([]byte, int, bool) {
|
||||
if len(pkt) < 20 {
|
||||
return nil, 0, false
|
||||
}
|
||||
switch pkt[0] >> 4 {
|
||||
case 4:
|
||||
if pkt[9] != wantProto {
|
||||
return nil, 0, false
|
||||
}
|
||||
trimmed, ok := fk.parseIPv4Prologue(pkt)
|
||||
return trimmed, 20, ok
|
||||
case 6:
|
||||
if len(pkt) < 40 {
|
||||
return nil, 0, false
|
||||
}
|
||||
if pkt[6] != wantProto {
|
||||
return nil, 0, false
|
||||
}
|
||||
trimmed, ok := fk.parseIPv6Prologue(pkt)
|
||||
return trimmed, 40, ok
|
||||
}
|
||||
return nil, 0, false
|
||||
}
|
||||
|
||||
// parseBase extracts the flow key and IP/TCP offsets for any TCP packet, admissible for
|
||||
// coalescing or not. Returns false for non-TCP or malformed input.
|
||||
func (p *parsedTCP) parseBase(pkt []byte) bool {
|
||||
trimmed, ipHdrLen, ok := p.fk.parseIPPrologue(pkt, ipProtoTCP)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return p.parseTail(trimmed, ipHdrLen)
|
||||
}
|
||||
|
||||
// parseBase extracts the flow key and IP/UDP offsets for a UDP packet.
|
||||
func (p *parsedUDP) parseBase(pkt []byte) bool {
|
||||
trimmed, ipHdrLen, ok := p.fk.parseIPPrologue(pkt, ipProtoUDP)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return p.parseTail(trimmed, ipHdrLen)
|
||||
}
|
||||
|
||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
||||
func (c *TCPCoalescer) Commit(pkt []byte) error {
|
||||
var info parsedTCP
|
||||
if !info.parseBase(pkt) {
|
||||
// Unparseable: flow key unknown, seal everything so later data cannot emit ahead of it.
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(pkt)
|
||||
return nil
|
||||
}
|
||||
return c.commitParsed(pkt, &info)
|
||||
}
|
||||
|
||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
||||
func (c *UDPCoalescer) Commit(pkt []byte) error {
|
||||
var info parsedUDP
|
||||
if !info.parseBase(pkt) {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(pkt)
|
||||
return nil
|
||||
}
|
||||
return c.commitParsed(pkt, &info)
|
||||
}
|
||||
@@ -1,121 +1,119 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"cmp"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"slices"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
)
|
||||
|
||||
// MultiCoalescer stages plaintext packets with their (epoch, counter) sort keys and, at Flush,
|
||||
// replays them in sender-transmission order into lane-specific batchers selected by L4 protocol.
|
||||
// MultiCoalescer fans plaintext packets out to lane-specific batchers based
|
||||
// on the IP/L4 protocol of the packet, sharing a single Reserve arena
|
||||
// across lanes so the caller's allocation pattern is unchanged.
|
||||
//
|
||||
// Sorting before dispatch keeps the ordering story simple: each lane consumes packets in
|
||||
// transmission order, builds slots in that order, and emits them in creation order. Wire reorder
|
||||
// inside a flush batch is repaired here, before it can fragment a lane's coalesce chains, so the
|
||||
// lanes carry no reorder-repair machinery.
|
||||
// Lanes are processed independently: the TCP coalescer only sees TCP, the
|
||||
// UDP coalescer only sees UDP, and the passthrough lane handles everything
|
||||
// else. Per-flow arrival order is preserved because a single 5-tuple only
|
||||
// ever lands in one lane and each lane preserves its own slot order.
|
||||
//
|
||||
// The contract is per-tunnel transmission order within each lane, with two exceptions: a pure TCP
|
||||
// ACK may be overtaken by later same-flow data (it does not close the flow's open chain; a late
|
||||
// ACK is just a stale ACK), and an unparseable shape seals every open chain in its lane (its flow
|
||||
// is unknown) and rides the lane as an in-lane verbatim, still in transmission order. Routing
|
||||
// follows the flow: a flow's non-coalesceable shapes ride its protocol lane rather than falling
|
||||
// to the later-flushed pt lane.
|
||||
//
|
||||
// Cross-lane order (TCP vs UDP vs everything else) is not preserved.
|
||||
// Cross-lane order is NOT preserved across the TCP/UDP/passthrough split.
|
||||
// This is acceptable because the carrier-side recvmmsg path already
|
||||
// stable-sorts by (peer, message counter) before delivering plaintext
|
||||
// here, so replay-window invariants are unaffected, and apps observe
|
||||
// correct per-flow ordering — which is all the IP layer guarantees anyway.
|
||||
// Do not "fix" this by interleaving lane outputs at flush time; that
|
||||
// negates the entire point of coalescing (each lane needs to see runs of
|
||||
// adjacent same-flow packets to coalesce them).
|
||||
type MultiCoalescer struct {
|
||||
tcp *TCPCoalescer
|
||||
udp *UDPCoalescer
|
||||
pt *Passthrough
|
||||
|
||||
// staged holds this batch's packets and sort keys until Flush. Borrowed: the caller keeps
|
||||
// each pkt alive until Flush returns.
|
||||
staged []stagedPacket
|
||||
// arena is owned by the Multi: lanes get only its Reserve (nil Resetter)
|
||||
// and Flush resets it exactly once after every lane has drained.
|
||||
arena *Arena
|
||||
}
|
||||
|
||||
// stagedPacket carries the scalars dispatch needs from the firewall's ParsedPacket, copied by
|
||||
// value: pp is reused by the caller per packet and must not be retained past Commit.
|
||||
type stagedPacket struct {
|
||||
pkt []byte
|
||||
key SortKey
|
||||
proto byte
|
||||
fragAny bool
|
||||
ipHdrLen uint16
|
||||
}
|
||||
// DefaultMultiArenaCap is the recommended arena capacity for a Multi-lane
|
||||
// batcher: 64 slots × 65535 bytes ≈ 4 MiB, enough to hold one recvmmsg
|
||||
// burst worth of MTU-sized packets without the arena growing.
|
||||
const DefaultMultiArenaCap = initialSlots * 65535
|
||||
|
||||
// NewMultiCoalescer builds a multi-lane batcher over w, based on available protocol support. The
|
||||
// staging sort applies even when no GSO lane is available: passthrough-only platforms still get
|
||||
// transmission-order repair.
|
||||
func NewMultiCoalescer(w io.Writer, l *slog.Logger) *MultiCoalescer {
|
||||
// NewMultiCoalescer builds a multi-lane batcher. tcpEnabled lets the caller
|
||||
// opt out of TCP coalescing (e.g. when the queue can't do TSO); udpEnabled
|
||||
// likewise gates UDP coalescing (only enable when USO was negotiated).
|
||||
// Either lane disabled redirects its traffic into the passthrough lane.
|
||||
// arena is the single backing slab shared across every lane; the caller
|
||||
// pre-sizes it via NewArena so the hot path never allocates.
|
||||
func NewMultiCoalescer(w io.Writer, l *slog.Logger, arena *Arena, tcpEnabled, udpEnabled bool) *MultiCoalescer {
|
||||
m := &MultiCoalescer{
|
||||
pt: NewPassthrough(w),
|
||||
staged: make([]stagedPacket, 0, initialSlots),
|
||||
pt: NewPassthrough(w, arena.Reserve, nil),
|
||||
arena: arena,
|
||||
}
|
||||
if tcpEnabled {
|
||||
m.tcp = NewTCPCoalescer(w, l, arena.Reserve, nil)
|
||||
}
|
||||
if udpEnabled {
|
||||
m.udp = NewUDPCoalescer(w, arena.Reserve, nil)
|
||||
}
|
||||
m.tcp = NewTCPCoalescer(w, l)
|
||||
m.udp = NewUDPCoalescer(w)
|
||||
return m
|
||||
}
|
||||
|
||||
// Commit stages pkt for the next Flush; dispatch is deferred so it runs on packets already in
|
||||
// transmission order. key carries the packet's tunnel epoch and message counter. pkt is borrowed:
|
||||
// the caller must keep it valid until the next Flush and not re-use it, and Flush may patch a
|
||||
// coalesced packet's headers in place. pp is the firewall's parse of pkt and is borrowed only
|
||||
// for this call, so the fields dispatch needs are copied here.
|
||||
func (m *MultiCoalescer) Commit(pkt []byte, key SortKey, pp *firewall.ParsedPacket) error {
|
||||
m.staged = append(m.staged, stagedPacket{
|
||||
pkt: pkt,
|
||||
key: key,
|
||||
proto: pp.Protocol,
|
||||
fragAny: pp.FragAny,
|
||||
ipHdrLen: uint16(pp.IPHdrLen),
|
||||
})
|
||||
return nil
|
||||
func (m *MultiCoalescer) Reserve(sz int) []byte {
|
||||
return m.arena.Reserve(sz)
|
||||
}
|
||||
|
||||
// compareStaged orders staged packets by (epoch, counter)
|
||||
func compareStaged(a, b stagedPacket) int {
|
||||
if c := cmp.Compare(a.key.Epoch, b.key.Epoch); c != 0 {
|
||||
return c
|
||||
// Commit dispatches pkt to the appropriate lane based on IP version + L4
|
||||
// proto. Borrowed slice contract is identical to the single-lane batchers,
|
||||
// pkt must remain valid until the next Flush.
|
||||
//
|
||||
// On the success path the IP/TCP-or-UDP parse happens here once and the
|
||||
// parsed struct is handed to the lane via commitParsed so the lane doesn't
|
||||
// re-walk the header.
|
||||
func (m *MultiCoalescer) Commit(pkt []byte) error {
|
||||
if len(pkt) < 20 {
|
||||
return m.pt.Commit(pkt)
|
||||
}
|
||||
return cmp.Compare(a.key.Counter, b.key.Counter)
|
||||
}
|
||||
|
||||
// dispatch routes one staged packet to its protocol lane (see commitStaged), or to the verbatim
|
||||
// passthrough when the lane has no GSO support.
|
||||
func (m *MultiCoalescer) dispatch(sp stagedPacket) error {
|
||||
switch sp.proto {
|
||||
v := pkt[0] >> 4
|
||||
var proto byte
|
||||
switch v {
|
||||
case 4:
|
||||
proto = pkt[9]
|
||||
case 6:
|
||||
if len(pkt) < 40 {
|
||||
return m.pt.Commit(pkt)
|
||||
}
|
||||
proto = pkt[6]
|
||||
default:
|
||||
return m.pt.Commit(pkt)
|
||||
}
|
||||
switch proto {
|
||||
case ipProtoTCP:
|
||||
if m.tcp != nil {
|
||||
return m.tcp.commitStaged(sp)
|
||||
info, ok := parseTCPBase(pkt)
|
||||
if !ok {
|
||||
// Malformed/unsupported TCP shape (IP options, fragments, ...).
|
||||
// Handle this via passthrough support in the TCP coalescer, to attempt to preserve flow order.
|
||||
m.tcp.addPassthrough(pkt)
|
||||
return nil
|
||||
}
|
||||
return m.tcp.commitParsed(pkt, info)
|
||||
}
|
||||
case ipProtoUDP:
|
||||
if m.udp != nil {
|
||||
return m.udp.commitStaged(sp)
|
||||
info, ok := parseUDP(pkt)
|
||||
if !ok {
|
||||
m.udp.addPassthrough(pkt) //we could also m.pt.Commit() here I guess?
|
||||
return nil
|
||||
}
|
||||
return m.udp.commitParsed(pkt, info)
|
||||
}
|
||||
}
|
||||
return m.pt.enqueue(sp.pkt)
|
||||
return m.pt.Commit(pkt)
|
||||
}
|
||||
|
||||
// Flush sorts the staged batch into transmission order, replays it into the lanes, then flushes each lane.
|
||||
// Drains everything and returns the joined errors; one bad packet does not hold up the rest.
|
||||
// After Flush returns, committed payload slices may be recycled.
|
||||
// Flush drains every lane in a fixed order, then resets the shared arena once.
|
||||
// A lane error doesn't stop the remaining lanes; the joined errors are returned.
|
||||
func (m *MultiCoalescer) Flush() error {
|
||||
// Arrival order is already almost sorted (reorder is the exception), which pdqsort detects
|
||||
// and handles in near-linear time.
|
||||
slices.SortFunc(m.staged, compareStaged)
|
||||
|
||||
var errs []error
|
||||
for _, sp := range m.staged {
|
||||
if err := m.dispatch(sp); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
clear(m.staged) // drop borrowed pkt refs
|
||||
m.staged = m.staged[:0]
|
||||
|
||||
if m.tcp != nil {
|
||||
if err := m.tcp.Flush(); err != nil {
|
||||
errs = append(errs, err)
|
||||
@@ -129,5 +127,6 @@ func (m *MultiCoalescer) Flush() error {
|
||||
if err := m.pt.Flush(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
m.arena.Reset()
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
@@ -1,39 +1,17 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/test"
|
||||
)
|
||||
|
||||
// keySeq hands out SortKeys with ascending counters in a fixed epoch, for
|
||||
// tests where commit order IS transmission order.
|
||||
type keySeq struct {
|
||||
epoch, counter uint64
|
||||
}
|
||||
|
||||
func (k *keySeq) next() SortKey {
|
||||
k.counter++
|
||||
return SortKey{Epoch: k.epoch, Counter: k.counter}
|
||||
}
|
||||
|
||||
// newTestMultiCoalescer builds a batcher over w.
|
||||
func newTestMultiCoalescer(tb testing.TB, w io.Writer) *MultiCoalescer {
|
||||
tb.Helper()
|
||||
return NewMultiCoalescer(w, test.NewLogger())
|
||||
}
|
||||
|
||||
// TestMultiCoalescerRoutesByProto confirms TCP/UDP/other land in the right
|
||||
// lane: TCP and UDP get coalesced when their lanes are enabled, anything
|
||||
// else (ICMP here) falls through to plain Write.
|
||||
func TestMultiCoalescerRoutesByProto(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, true)
|
||||
|
||||
tcpPay := make([]byte, 1200)
|
||||
udpPay := make([]byte, 1200)
|
||||
@@ -43,19 +21,19 @@ func TestMultiCoalescerRoutesByProto(t *testing.T) {
|
||||
icmp[3] = 28
|
||||
icmp[9] = 1
|
||||
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(1000, tcpAck, tcpPay))); err != nil {
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(2200, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(2200, tcpAck, tcpPay))); err != nil {
|
||||
if err := m.Commit(buildTCPv4(2200, tcpAck, tcpPay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil {
|
||||
if err := m.Commit(buildUDPv4(2000, 53, udpPay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil {
|
||||
if err := m.Commit(buildUDPv4(2000, 53, udpPay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(icmp, k.next(), testPP(icmp)); err != nil {
|
||||
if err := m.Commit(icmp); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
@@ -70,162 +48,17 @@ func TestMultiCoalescerRoutesByProto(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerRestoresTransmissionOrder is the core staging-sort
|
||||
// property: packets committed out of counter order (wire reorder inside one
|
||||
// flush batch) are replayed into the lanes in transmission order, so the
|
||||
// reorder never fragments the coalesce chain — one superpacket, in seq
|
||||
// order, exactly as if the wire had never reordered. The retransmit shape
|
||||
// falls out of the same key: a retransmit carries a lower seq but a HIGHER
|
||||
// counter (it was encrypted later), so it emits after the data it trails.
|
||||
func TestMultiCoalescerRestoresTransmissionOrder(t *testing.T) {
|
||||
// TestMultiCoalescerDisabledUDPFallsThrough verifies that when the UDP lane
|
||||
// is disabled (e.g. kernel doesn't support USO), UDP packets still reach
|
||||
// the kernel via the passthrough lane rather than being lost.
|
||||
func TestMultiCoalescerDisabledUDPFallsThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, false) // TSO on, USO off
|
||||
|
||||
// Transmission order: seq 1000 (c1), 2200 (c2), 3400 (c3).
|
||||
// Arrival order: 3400, 1000, 2200.
|
||||
if err := m.Commit(buildTCPv4(3400, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4(3400, tcpAck, pay))); err != nil {
|
||||
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 1, Counter: 1}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(2200, tcpAck, pay), SortKey{Epoch: 1, Counter: 2}, testPP(buildTCPv4(2200, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 1 || len(w.writes) != 0 {
|
||||
t.Fatalf("want 1 gso write (unfragmented chain), got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
g := w.gsoWrites[0]
|
||||
if len(g.pays) != 3 {
|
||||
t.Fatalf("segs=%d want 3", len(g.pays))
|
||||
}
|
||||
const ipHdrLen = 20
|
||||
if seedSeq := binary.BigEndian.Uint32(g.hdr[ipHdrLen+4 : ipHdrLen+8]); seedSeq != 1000 {
|
||||
t.Errorf("seed seq=%d want 1000", seedSeq)
|
||||
}
|
||||
|
||||
// Retransmit: seq 1000 again but counter 4 — sorts after seq 4600 (c3).
|
||||
w.writes, w.gsoWrites, w.order = nil, nil, nil
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 1, Counter: 4}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(4600, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4(4600, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 2 {
|
||||
t.Fatalf("want 2 plain writes, got %d (gso=%d)", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
first := binary.BigEndian.Uint32(w.writes[0][24:28])
|
||||
second := binary.BigEndian.Uint32(w.writes[1][24:28])
|
||||
if first != 4600 || second != 1000 {
|
||||
t.Fatalf("emission (%d, %d), want (4600, 1000): retransmit must not overtake in-flight data", first, second)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerRestoresOrderAcrossFlows scrambles two interleaved flows;
|
||||
// the staging sort must repair each flow into one superpacket without any
|
||||
// cross-flow contamination.
|
||||
func TestMultiCoalescerRestoresOrderAcrossFlows(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
|
||||
// Transmission: A.100 (c1), B.500 (c2), A.1300 (c3), B.1700 (c4).
|
||||
// Arrival: A.1300, B.1700, A.100, B.500.
|
||||
if err := m.Commit(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay), SortKey{Epoch: 1, Counter: 4}, testPP(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay), SortKey{Epoch: 1, Counter: 1}, testPP(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay), SortKey{Epoch: 1, Counter: 2}, testPP(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 gso writes (one per flow), got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
for i, g := range w.gsoWrites {
|
||||
if len(g.pays) != 2 {
|
||||
t.Errorf("gso[%d] segs=%d want 2", i, len(g.pays))
|
||||
}
|
||||
const ipHdrLen = 20
|
||||
seedSeq := binary.BigEndian.Uint32(g.hdr[ipHdrLen+4 : ipHdrLen+8])
|
||||
sport := binary.BigEndian.Uint16(g.hdr[ipHdrLen : ipHdrLen+2])
|
||||
switch sport {
|
||||
case 1000:
|
||||
if seedSeq != 100 {
|
||||
t.Errorf("flow A seed seq=%d want 100", seedSeq)
|
||||
}
|
||||
case 3000:
|
||||
if seedSeq != 500 {
|
||||
t.Errorf("flow B seed seq=%d want 500", seedSeq)
|
||||
}
|
||||
default:
|
||||
t.Errorf("unexpected sport %d", sport)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerEpochOrdersAcrossRehandshake: a re-handshake replaces
|
||||
// the tunnel, and the replacement's counter space starts near zero — raw
|
||||
// counter order would emit the new tunnel's packets first while the old
|
||||
// tunnel's backlog is still arriving. The epoch key must dominate:
|
||||
// everything from the old tunnel emits before anything from the new one.
|
||||
func TestMultiCoalescerEpochOrdersAcrossRehandshake(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
|
||||
// New session's first data arrives before the old session's last data.
|
||||
if err := m.Commit(buildTCPv4(2200, tcpAck, pay), SortKey{Epoch: 8, Counter: 1}, testPP(buildTCPv4(2200, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 7, Counter: 9_000_000}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Same flow, contiguous seq, identical headers: after the epoch sort the
|
||||
// two segments append into one superpacket seeded by the OLD session's
|
||||
// packet.
|
||||
if len(w.gsoWrites) != 1 {
|
||||
t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
const ipHdrLen = 20
|
||||
if seedSeq := binary.BigEndian.Uint32(w.gsoWrites[0].hdr[ipHdrLen+4 : ipHdrLen+8]); seedSeq != 1000 {
|
||||
t.Errorf("seed seq=%d want 1000 (old session first)", seedSeq)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerNoUSOFallsThrough verifies that on a queue without USO
|
||||
// (older kernel: TSO but no GSO_UDP_L4) the UDP lane never comes up and UDP
|
||||
// packets still reach the kernel via verbatim rather than being lost.
|
||||
func TestMultiCoalescerNoUSOFallsThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true, noUSO: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
if m.udp != nil {
|
||||
t.Fatal("UDP lane must not come up without USO")
|
||||
}
|
||||
|
||||
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv4(1000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv4(1000, 53, make([]byte, 800)))); err != nil {
|
||||
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
@@ -239,164 +72,16 @@ func TestMultiCoalescerNoUSOFallsThrough(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerNoOffloadsStillSorts covers a queue that can't offload
|
||||
// anything. Both lane constructors refuse, so every packet rides the
|
||||
// verbatim lane — but the staging sort still applies, so emission follows
|
||||
// transmission order even without GSO.
|
||||
func TestMultiCoalescerNoOffloadsStillSorts(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: false}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
if m.tcp != nil || m.udp != nil {
|
||||
t.Fatal("no lane may come up without offloads")
|
||||
}
|
||||
pkts := [][]byte{
|
||||
buildTCPv4(1000, tcpAck, make([]byte, 1200)),
|
||||
buildUDPv4(1000, 53, make([]byte, 800)),
|
||||
buildTCPv4(2200, tcpAck, make([]byte, 1200)),
|
||||
}
|
||||
// Committed in reverse transmission order; keys carry the truth.
|
||||
for i := len(pkts) - 1; i >= 0; i-- {
|
||||
if err := m.Commit(pkts[i], SortKey{Epoch: 1, Counter: uint64(i + 1)}, testPP(pkts[i])); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 0 {
|
||||
t.Errorf("no GSO writes possible, got %d", len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes) != len(pkts) {
|
||||
t.Fatalf("want %d plain writes, got %d", len(pkts), len(w.writes))
|
||||
}
|
||||
// One lane for everything means the sorted order survives end to end.
|
||||
for i, want := range pkts {
|
||||
if !bytes.Equal(w.writes[i], want) {
|
||||
t.Errorf("write %d out of order or corrupt", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// buildUDPv6Fragment builds an IPv6 packet whose extension chain is a
|
||||
// single fragment header (NH=44) naming UDP as the terminal protocol —
|
||||
// a first fragment (offset 0, MF set) carrying the UDP header and a
|
||||
// partial payload.
|
||||
func buildUDPv6Fragment(sport, dport uint16, payload []byte) []byte {
|
||||
const ipHdrLen = 40
|
||||
const fragHdrLen = 8
|
||||
const udpHdrLen = 8
|
||||
total := ipHdrLen + fragHdrLen + udpHdrLen + len(payload)
|
||||
pkt := make([]byte, total)
|
||||
|
||||
pkt[0] = 0x60
|
||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(total-ipHdrLen))
|
||||
pkt[6] = 44 // fragment extension header
|
||||
pkt[7] = 64
|
||||
pkt[8] = 0xfe
|
||||
pkt[9] = 0x80
|
||||
pkt[23] = 1
|
||||
pkt[24] = 0xfe
|
||||
pkt[25] = 0x80
|
||||
pkt[39] = 2
|
||||
|
||||
pkt[40] = ipProtoUDP // fragment's next header
|
||||
binary.BigEndian.PutUint16(pkt[42:44], 0x0001) // offset 0, MF set
|
||||
binary.BigEndian.PutUint32(pkt[44:48], 0x1badf00) // identification
|
||||
|
||||
binary.BigEndian.PutUint16(pkt[48:50], sport)
|
||||
binary.BigEndian.PutUint16(pkt[50:52], dport)
|
||||
binary.BigEndian.PutUint16(pkt[52:54], uint16(udpHdrLen+len(payload)))
|
||||
copy(pkt[56:], payload)
|
||||
return pkt
|
||||
}
|
||||
|
||||
// TestMultiCoalescerIPv6FragmentStaysInLane locks in extension-header
|
||||
// routing: a fragment whose chain terminates in UDP must ride the UDP lane
|
||||
// as an in-lane verbatim — emitted ahead of later same-flow datagrams —
|
||||
// not the verbatim lane, which flushes after every coalescer lane and
|
||||
// would reorder it behind data that arrived after it.
|
||||
func TestMultiCoalescerIPv6FragmentStaysInLane(t *testing.T) {
|
||||
// TestMultiCoalescerDisabledTCPFallsThrough mirrors the TSO=off case.
|
||||
func TestMultiCoalescerDisabledTCPFallsThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
|
||||
if err := m.Commit(buildUDPv6Fragment(2000, 53, make([]byte, 512)), k.next(), testPP(buildUDPv6Fragment(2000, 53, make([]byte, 512)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 1 {
|
||||
t.Fatalf("want the fragment as 1 plain write, got %d", len(w.writes))
|
||||
}
|
||||
if len(w.gsoWrites) != 1 {
|
||||
t.Fatalf("want the two whole datagrams coalesced into 1 gso write, got %d", len(w.gsoWrites))
|
||||
}
|
||||
// Transmission order was fragment-then-data; same-lane routing must keep it.
|
||||
if w.order[0] != "write" {
|
||||
t.Fatalf("fragment must be emitted before later data (in-lane verbatim), order=%v", w.order)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerFragmentSealsUDPChains: an unparseable datagram
|
||||
// (fragment) seals every open UDP chain, so datagrams from before and after
|
||||
// it land in separate superpackets and the fragment holds its transmission-
|
||||
// order position between them.
|
||||
func TestMultiCoalescerFragmentSealsUDPChains(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6Fragment(2000, 53, make([]byte, 512)), k.next(), testPP(buildUDPv6Fragment(2000, 53, make([]byte, 512)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 gso writes (chains sealed around the fragment), got %d", len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes) != 1 {
|
||||
t.Fatalf("want the fragment as 1 plain write, got %d", len(w.writes))
|
||||
}
|
||||
want := []string{"gso", "write", "gso"}
|
||||
if len(w.order) != 3 || w.order[0] != want[0] || w.order[1] != want[1] || w.order[2] != want[2] {
|
||||
t.Fatalf("emission order = %v, want %v", w.order, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerNoTSOFallsThrough mirrors the no-TSO case.
|
||||
func TestMultiCoalescerNoTSOFallsThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true, noTSO: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
if m.tcp != nil {
|
||||
t.Fatal("TCP lane must not come up without TSO")
|
||||
}
|
||||
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), false, true) // TSO off, USO on
|
||||
|
||||
pay := make([]byte, 1200)
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), k.next(), testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(2200, tcpAck, pay), k.next(), testPP(buildTCPv4(2200, tcpAck, pay))); err != nil {
|
||||
if err := m.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
@@ -409,29 +94,3 @@ func TestMultiCoalescerNoTSOFallsThrough(t *testing.T) {
|
||||
t.Errorf("TCP must pass through as 2 plain writes, got %d", len(w.writes))
|
||||
}
|
||||
}
|
||||
|
||||
// testPP derives the ParsedPacket newPacket would produce for the packet
|
||||
// shapes the tests build: plain v4/v6, v4 with options or fragment bits set,
|
||||
// and the single-fragment-header v6 shape from buildUDPv6Fragment. Anything
|
||||
// unrecognizable stays zero (proto 0 routes to the passthrough lane).
|
||||
func testPP(pkt []byte) *firewall.ParsedPacket {
|
||||
pp := &firewall.ParsedPacket{}
|
||||
if len(pkt) < 20 {
|
||||
return pp
|
||||
}
|
||||
switch pkt[0] >> 4 {
|
||||
case 4:
|
||||
pp.Protocol = pkt[9]
|
||||
pp.IPHdrLen = int(pkt[0]&0x0f) * 4
|
||||
pp.FragAny = binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0
|
||||
case 6:
|
||||
pp.Protocol = pkt[6]
|
||||
pp.IPHdrLen = 40
|
||||
if pp.Protocol == 44 { // fragment extension header
|
||||
pp.Protocol = pkt[40]
|
||||
pp.IPHdrLen = 48
|
||||
pp.FragAny = true
|
||||
}
|
||||
}
|
||||
return pp
|
||||
}
|
||||
|
||||
@@ -2,29 +2,54 @@ package batch
|
||||
|
||||
import (
|
||||
"io"
|
||||
|
||||
"github.com/slackhq/nebula/udp"
|
||||
)
|
||||
|
||||
// Passthrough is MultiCoalescer's verbatim lane: no batching, packets are written at Flush in the
|
||||
// order enqueued.
|
||||
// Passthrough is a RxBatcher that doesn't batch anything, it just accumulates and then sends packets.
|
||||
type Passthrough struct {
|
||||
out io.Writer
|
||||
slots [][]byte
|
||||
out io.Writer
|
||||
slots [][]byte
|
||||
reserver Reserver
|
||||
resetter Resetter
|
||||
cursor int
|
||||
}
|
||||
|
||||
func NewPassthrough(w io.Writer) *Passthrough {
|
||||
const passthroughBaseNumSlots = 128
|
||||
|
||||
// DefaultPassthroughArenaCap is the recommended arena capacity for a
|
||||
// standalone Passthrough batcher: 128 slots × udp.MTU ≈ 1.1 MiB.
|
||||
const DefaultPassthroughArenaCap = passthroughBaseNumSlots * udp.MTU
|
||||
|
||||
func NewPassthrough(w io.Writer, reserver Reserver, resetter Resetter) *Passthrough {
|
||||
return &Passthrough{
|
||||
out: w,
|
||||
slots: make([][]byte, 0, 128),
|
||||
out: w,
|
||||
slots: make([][]byte, 0, passthroughBaseNumSlots),
|
||||
reserver: reserver,
|
||||
resetter: resetter,
|
||||
}
|
||||
}
|
||||
|
||||
// enqueue accepts one packet, already sorted into transmission order by dispatch.
|
||||
func (p *Passthrough) enqueue(pkt []byte) error {
|
||||
func (p *Passthrough) Reserve(sz int) []byte {
|
||||
return p.reserver(sz)
|
||||
}
|
||||
|
||||
func (p *Passthrough) Commit(pkt []byte) error {
|
||||
p.slots = append(p.slots, pkt)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Flush drains every queued packet and calls the configured Resetter
|
||||
func (p *Passthrough) Flush() error {
|
||||
firstErr := p.drain()
|
||||
if p.resetter != nil {
|
||||
p.resetter()
|
||||
}
|
||||
return firstErr
|
||||
}
|
||||
|
||||
// drain writes out every queued packet and clears the slot list.
|
||||
func (p *Passthrough) drain() error {
|
||||
var firstErr error
|
||||
for _, s := range p.slots {
|
||||
_, err := p.out.Write(s)
|
||||
|
||||
+432
-176
@@ -2,9 +2,12 @@ package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"slices"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
)
|
||||
@@ -20,19 +23,24 @@ const tcpCoalesceBufSize = 65535
|
||||
// superpacket. Keeping this well below the kernel's TSO ceiling bounds latency.
|
||||
const tcpCoalesceMaxSegs = 64
|
||||
|
||||
// coalesceSlot is one entry in the coalescer's ordered event queue. A verbatim slot holds a single
|
||||
// borrowed packet emitted as-is (pure ACK, non-admissible TCP, unparseable, or oversize seed); a
|
||||
// non-verbatim slot is an in-progress coalesced superpacket. payIovs are borrowed slices of the
|
||||
// caller's plaintext buffers; the caller must keep them alive until Flush.
|
||||
// tcpCoalesceHdrCap is the scratch space we copy a seed's IP+TCP header
|
||||
// into. IPv6 (40) + TCP with full options (60) = 100 bytes.
|
||||
const tcpCoalesceHdrCap = 100
|
||||
|
||||
// coalesceSlot is one entry in the coalescer's ordered event queue. When
|
||||
// passthrough is true the slot holds a single borrowed packet that must be
|
||||
// emitted verbatim (non-TCP, non-admissible TCP, or oversize seed). When
|
||||
// passthrough is false the slot is an in-progress coalesced superpacket:
|
||||
// hdrBuf is a mutable copy of the seed's IP+TCP header (we patch total
|
||||
// length and pseudo-header partial at flush), and payIovs are *borrowed*
|
||||
// slices from the caller's plaintext buffers — no payload is ever copied.
|
||||
// The caller (listenOut) must keep those buffers alive until Flush.
|
||||
type coalesceSlot struct {
|
||||
verbatim bool
|
||||
// rawPkt is borrowed: the whole packet for verbatim slots, the seed packet for coalesce
|
||||
// slots. A slot that never grows past one segment is emitted from rawPkt so its original
|
||||
// (already valid) L4 checksum ships DATA_VALID instead of making the kernel recompute it.
|
||||
// A multi-segment slot's superpacket header is rawPkt's, patched in place at flush.
|
||||
rawPkt []byte
|
||||
passthrough bool
|
||||
rawPkt []byte // borrowed when passthrough
|
||||
|
||||
fk flowKey
|
||||
hdrBuf [tcpCoalesceHdrCap]byte
|
||||
hdrLen int
|
||||
ipHdrLen int
|
||||
isV6 bool
|
||||
@@ -40,203 +48,216 @@ type coalesceSlot struct {
|
||||
numSeg int
|
||||
totalPay int
|
||||
nextSeq uint32
|
||||
payIovs [][]byte
|
||||
// psh closes the chain: set when the last-accepted segment had PSH or
|
||||
// was sub-gsoSize. No further appends after that.
|
||||
psh bool
|
||||
payIovs [][]byte
|
||||
}
|
||||
|
||||
// TCPCoalescer accumulates adjacent in-flow TCP data segments across multiple concurrent flows and
|
||||
// emits each flow's run as a single TSO superpacket via tio.GSOWriter. Input must be in sender
|
||||
// transmission order (MultiCoalescer sorts by (epoch, counter) before dispatch); slots are emitted
|
||||
// in creation order, so emission reproduces transmission order except for the pure-ACK case in
|
||||
// commitParsed. Owns no locks; one coalescer per TUN write queue.
|
||||
// TCPCoalescer accumulates adjacent in-flow TCP data segments across
|
||||
// multiple concurrent flows and emits each flow's run as a single TSO
|
||||
// superpacket via tio.GSOWriter. All output — coalesced or not — is
|
||||
// deferred until Flush so arrival order is preserved on the wire. Owns
|
||||
// no locks; one coalescer per TUN write queue.
|
||||
type TCPCoalescer struct {
|
||||
w tio.GSOWriter
|
||||
plainW io.Writer
|
||||
gsoW tio.GSOWriter // nil when the queue doesn't support TSO
|
||||
|
||||
// slots is the ordered event queue. Flush walks it once and emits each
|
||||
// entry as either a WriteGSO (coalesced) or a w.Write (verbatim).
|
||||
// entry as either a WriteGSO (coalesced) or a plainW.Write (passthrough).
|
||||
slots []*coalesceSlot
|
||||
// openSlots maps a flow key to its open slot so new segments can extend an in-progress
|
||||
// superpacket in O(1). Removal is what closes a chain: on PSH or a short last segment, on a
|
||||
// non-admissible packet for the flow, or in Flush.
|
||||
// openSlots maps a flow key to its most recent non-sealed slot, so new
|
||||
// segments can extend an in-progress superpacket in O(1). Slots are
|
||||
// removed from this map when they close (PSH or short-last-segment),
|
||||
// when a non-admissible packet for that flow arrives, or in Flush.
|
||||
openSlots map[flowKey]*coalesceSlot
|
||||
// lastSlot caches the most recently touched open slot. Bulk traffic
|
||||
// arrives in same-flow runs (single-flow steady state, or GRO bursts
|
||||
// under multi-flow), so comparing the incoming key against the cached
|
||||
// slot's own fk lets the hot path skip the map lookup (and the aeshash
|
||||
// of a 38-byte key) for the length of each run.
|
||||
// lastSlot caches the most recently touched open slot. Steady-state
|
||||
// bulk traffic is dominated by a single flow, so comparing the
|
||||
// incoming key against the cached slot's own fk lets the hot path
|
||||
// skip the map lookup (and the aeshash of a 38-byte key) entirely.
|
||||
// Kept in lockstep with openSlots: nil whenever the slot it pointed
|
||||
// at is removed.
|
||||
// at is removed/sealed.
|
||||
lastSlot *coalesceSlot
|
||||
pool []*coalesceSlot // free list for reuse
|
||||
reserver Reserver
|
||||
resetter Resetter
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
// NewTCPCoalescer wraps w, returning nil if w can't accept GSO_TCP writes.
|
||||
func NewTCPCoalescer(w io.Writer, l *slog.Logger) *TCPCoalescer {
|
||||
gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
return &TCPCoalescer{
|
||||
w: gw,
|
||||
func NewTCPCoalescer(w io.Writer, l *slog.Logger, reserver Reserver, resetter Resetter) *TCPCoalescer {
|
||||
c := &TCPCoalescer{
|
||||
plainW: w,
|
||||
slots: make([]*coalesceSlot, 0, initialSlots),
|
||||
openSlots: make(map[flowKey]*coalesceSlot, initialSlots),
|
||||
pool: make([]*coalesceSlot, 0, initialSlots),
|
||||
reserver: reserver,
|
||||
resetter: resetter,
|
||||
l: l,
|
||||
}
|
||||
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP); ok {
|
||||
c.gsoW = gw
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
// parsedTCP holds the fields extracted from a single parse so later steps
|
||||
// (admission, slot lookup, canAppend) don't re-walk the header.
|
||||
type parsedTCP struct {
|
||||
fk flowKey
|
||||
ipHdrLen int
|
||||
hdrLen int
|
||||
payLen int
|
||||
seq uint32
|
||||
flags byte
|
||||
fk flowKey
|
||||
ipHdrLen int
|
||||
tcpHdrLen int
|
||||
hdrLen int
|
||||
payLen int
|
||||
seq uint32
|
||||
flags byte
|
||||
}
|
||||
|
||||
// parseAt extracts the flow key and IP/TCP offsets for a packet the dispatcher already knows is
|
||||
// TCP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on
|
||||
// entry and is filled in place; see flowKey.parseIPAt for why. Returns false for malformed input
|
||||
// or any shape that must not coalesce (IPv4 options/fragmentation, IPv6 extension headers).
|
||||
func (p *parsedTCP) parseAt(pkt []byte, ipHdrLen int) bool {
|
||||
trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen)
|
||||
// parseTCPBase extracts the flow key and IP/TCP offsets for any TCP packet,
|
||||
// regardless of whether it's admissible for coalescing. Returns ok=false
|
||||
// for non-TCP or malformed input.
|
||||
// Accepts IPv4 (no options or fragmentation) and IPv6 (no extension headers).
|
||||
func parseTCPBase(pkt []byte) (parsedTCP, bool) {
|
||||
var p parsedTCP
|
||||
ip, ok := parseIPPrologue(pkt, ipProtoTCP)
|
||||
if !ok {
|
||||
return false
|
||||
return p, false
|
||||
}
|
||||
return p.parseTail(trimmed, ipHdrLen)
|
||||
}
|
||||
pkt = ip.pkt
|
||||
p.fk = ip.fk
|
||||
p.ipHdrLen = ip.ipHdrLen
|
||||
|
||||
// parseTail layers the TCP-header parse on a validated IP prologue. pkt is the trimmed packet;
|
||||
// fk's addresses are already filled.
|
||||
func (p *parsedTCP) parseTail(pkt []byte, ipHdrLen int) bool {
|
||||
if len(pkt) < ipHdrLen+20 {
|
||||
return false
|
||||
if len(pkt) < p.ipHdrLen+20 {
|
||||
return p, false
|
||||
}
|
||||
tcpOff := int(pkt[ipHdrLen+12]>>4) * 4
|
||||
tcpOff := int(pkt[p.ipHdrLen+12]>>4) * 4
|
||||
if tcpOff < 20 || tcpOff > 60 {
|
||||
return false
|
||||
return p, false
|
||||
}
|
||||
if len(pkt) < ipHdrLen+tcpOff {
|
||||
return false
|
||||
if len(pkt) < p.ipHdrLen+tcpOff {
|
||||
return p, false
|
||||
}
|
||||
p.ipHdrLen = ipHdrLen
|
||||
p.hdrLen = ipHdrLen + tcpOff
|
||||
p.tcpHdrLen = tcpOff
|
||||
p.hdrLen = p.ipHdrLen + tcpOff
|
||||
p.payLen = len(pkt) - p.hdrLen
|
||||
p.fk.sport = binary.BigEndian.Uint16(pkt[ipHdrLen : ipHdrLen+2])
|
||||
p.fk.dport = binary.BigEndian.Uint16(pkt[ipHdrLen+2 : ipHdrLen+4])
|
||||
p.seq = binary.BigEndian.Uint32(pkt[ipHdrLen+4 : ipHdrLen+8])
|
||||
p.flags = pkt[ipHdrLen+13]
|
||||
return true
|
||||
p.seq = binary.BigEndian.Uint32(pkt[p.ipHdrLen+4 : p.ipHdrLen+8])
|
||||
p.flags = pkt[p.ipHdrLen+13]
|
||||
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2])
|
||||
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
|
||||
return p, true
|
||||
}
|
||||
|
||||
// TCP flag bits (byte 13 of the TCP header). Only the bits the coalescer consults are named;
|
||||
// FIN/SYN/RST/URG/CWR are rejected by the negative mask in commitParsed.
|
||||
// TCP flag bits (byte 13 of the TCP header). Only the bits actually consulted
|
||||
// by the coalescer are named; FIN/SYN/RST/URG/CWR are rejected via the
|
||||
// negative mask in coalesceable, not by name.
|
||||
const (
|
||||
tcpFlagPsh = 0x08
|
||||
tcpFlagAck = 0x10
|
||||
tcpFlagEce = 0x40
|
||||
)
|
||||
|
||||
// sealAllOpen closes every open coalesce chain. Called for unparseable packets: the flow key is
|
||||
// unknown, so any open chain could otherwise absorb later data and emit it ahead of this packet.
|
||||
func (c *TCPCoalescer) sealAllOpen() {
|
||||
clear(c.openSlots)
|
||||
c.lastSlot = nil
|
||||
// coalesceable reports whether a parsed TCP segment is eligible for
|
||||
// coalescing. Accepts ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE with a
|
||||
// non-empty payload. CWR is excluded because it marks a one-shot
|
||||
// congestion-window-reduced transition the receiver must observe at a
|
||||
// segment boundary.
|
||||
func (p parsedTCP) coalesceable() bool {
|
||||
if p.flags&tcpFlagAck == 0 {
|
||||
return false
|
||||
}
|
||||
if p.flags&^(tcpFlagAck|tcpFlagPsh|tcpFlagEce) != 0 {
|
||||
return false
|
||||
}
|
||||
return p.payLen > 0
|
||||
}
|
||||
|
||||
// sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips
|
||||
// hashing the 38-byte key when no chains are open (e.g. ack-dominant queues).
|
||||
func (c *TCPCoalescer) sealFlow(fk flowKey) {
|
||||
if len(c.openSlots) == 0 {
|
||||
return
|
||||
}
|
||||
if last := c.lastSlot; last != nil && last.fk == fk {
|
||||
c.lastSlot = nil
|
||||
}
|
||||
delete(c.openSlots, fk)
|
||||
func (c *TCPCoalescer) Reserve(sz int) []byte {
|
||||
return c.reserver(sz)
|
||||
}
|
||||
|
||||
// commitStaged commits one staged packet dispatch routed to this lane. A shape the lane cannot
|
||||
// coalesce (any fragmentation, unparseable header) seals every open chain
|
||||
// and rides the lane as an in-lane verbatim, still in transmission order.
|
||||
func (c *TCPCoalescer) commitStaged(sp stagedPacket) error {
|
||||
if sp.fragAny {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(sp.pkt)
|
||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
||||
func (c *TCPCoalescer) Commit(pkt []byte) error {
|
||||
if c.gsoW == nil {
|
||||
c.addPassthrough(pkt)
|
||||
return nil
|
||||
}
|
||||
var info parsedTCP
|
||||
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(sp.pkt)
|
||||
info, ok := parseTCPBase(pkt)
|
||||
if !ok {
|
||||
c.addPassthrough(pkt)
|
||||
return nil
|
||||
}
|
||||
return c.commitParsed(sp.pkt, &info)
|
||||
return c.commitParsed(pkt, info)
|
||||
}
|
||||
|
||||
// commitParsed commits one parsed TCP packet. The caller (dispatch, via parseAt) supplies a
|
||||
// valid parse so the header is not re-walked here.
|
||||
func (c *TCPCoalescer) commitParsed(pkt []byte, info *parsedTCP) error {
|
||||
// Admission: only ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE may ride a coalesce chain. CWR marks a
|
||||
// one-shot congestion transition the receiver must observe at a segment boundary. NB: AccECN
|
||||
// reuses CWR as ACE counter bits; revisit this check if inner hosts adopt AccECN.
|
||||
if info.flags&tcpFlagAck == 0 || info.flags&^(tcpFlagAck|tcpFlagPsh|tcpFlagEce) != 0 {
|
||||
// SYN/FIN/RST/URG/CWR must be observed in sequence. Seal the flow's open slot so later
|
||||
// in-flow packets cannot extend it and emit ahead of this verbatim.
|
||||
c.sealFlow(info.fk)
|
||||
c.addVerbatim(pkt)
|
||||
// commitParsed is the post-parse half of Commit. The caller must have
|
||||
// already verified parseTCPBase succeeded (info is a valid TCP parse).
|
||||
// Used by MultiCoalescer.Commit to avoid re-walking the IP/TCP header
|
||||
// after the dispatcher has already done so.
|
||||
func (c *TCPCoalescer) commitParsed(pkt []byte, info parsedTCP) error {
|
||||
if c.gsoW == nil {
|
||||
c.addPassthrough(pkt)
|
||||
return nil
|
||||
}
|
||||
if info.payLen == 0 {
|
||||
// Pure ACK: no ordering obligation toward the flow's data. Delivering it after
|
||||
// later-transmitted data only makes it a stale ACK, which receivers ignore. Not sealing
|
||||
// keeps a bidirectional flow's data run coalescing across interleaved peer ACKs, matching
|
||||
// kernel GRO. This is the only place emission deviates from transmission order.
|
||||
c.addVerbatim(pkt)
|
||||
if !info.coalesceable() {
|
||||
// TCP but not admissible (SYN/FIN/RST/URG/CWR or zero-payload).
|
||||
// Seal this flow's open slot so later in-flow packets don't extend
|
||||
// it and accidentally reorder past this passthrough.
|
||||
if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||
c.lastSlot = nil
|
||||
}
|
||||
delete(c.openSlots, info.fk)
|
||||
c.addPassthrough(pkt)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Cached-slot fast path. Arrival isn't per-packet interleaved even with
|
||||
// many flows: wire-side GRO delivers runs of same-flow packets
|
||||
// (deliverSegments splits a superdatagram into up to 64), so the cache
|
||||
// hits for the length of each run and a miss costs one fk compare
|
||||
// before the map lookup carries the weight.
|
||||
// Single-flow fast path: with only one open flow the cache hits every
|
||||
// packet, and len(openSlots)==1 lets us skip the 38-byte fk compare
|
||||
// when there are multiple flows in flight (where the hit rate would
|
||||
// be ~0 and the compare is pure overhead).
|
||||
var open *coalesceSlot
|
||||
if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||
if last := c.lastSlot; last != nil && len(c.openSlots) == 1 && last.fk == info.fk {
|
||||
open = last
|
||||
} else {
|
||||
open = c.openSlots[info.fk]
|
||||
}
|
||||
if open != nil {
|
||||
if c.canAppend(open, pkt, info) {
|
||||
if c.appendPayload(open, pkt, info) {
|
||||
// Chain closed (PSH or short segment): stop extending it.
|
||||
c.sealFlow(info.fk)
|
||||
c.appendPayload(open, pkt, info)
|
||||
if open.psh {
|
||||
delete(c.openSlots, info.fk)
|
||||
c.lastSlot = nil
|
||||
} else {
|
||||
c.lastSlot = open
|
||||
}
|
||||
return nil
|
||||
}
|
||||
// Can't extend (seq gap from upstream loss, header change, or a full
|
||||
// chain): evict it from openSlots and fall through to seed a fresh slot.
|
||||
c.sealFlow(info.fk)
|
||||
// Can't extend — seal it and fall through to seed a fresh slot.
|
||||
delete(c.openSlots, info.fk)
|
||||
if c.lastSlot == open {
|
||||
c.lastSlot = nil
|
||||
}
|
||||
}
|
||||
c.seed(pkt, info)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Flush emits every queued event in (per-flow) seq order.
|
||||
func (c *TCPCoalescer) Flush() error {
|
||||
first := c.drain()
|
||||
if c.resetter != nil {
|
||||
c.resetter()
|
||||
}
|
||||
return first
|
||||
}
|
||||
|
||||
// drain emits every queued slot (reordering/merging coalesced runs first)
|
||||
// and clears the slot state.
|
||||
func (c *TCPCoalescer) drain() error {
|
||||
c.reorderForFlush()
|
||||
var first error
|
||||
for _, s := range c.slots {
|
||||
var err error
|
||||
if s.verbatim || s.numSeg == 1 {
|
||||
// A slot that never grew is byte-identical to its seed packet; ship the original so
|
||||
// its valid checksum rides the DATA_VALID path instead of a kernel software csum.
|
||||
// rawPkt is only mutated once numSeg >= 2 (PSH propagate, flush patches), so it is
|
||||
// pristine here.
|
||||
_, err = c.w.Write(s.rawPkt)
|
||||
if s.passthrough {
|
||||
_, err = c.plainW.Write(s.rawPkt)
|
||||
} else {
|
||||
err = c.flushSlot(s)
|
||||
}
|
||||
@@ -253,27 +274,23 @@ func (c *TCPCoalescer) Flush() error {
|
||||
return first
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) addVerbatim(pkt []byte) {
|
||||
func (c *TCPCoalescer) addPassthrough(pkt []byte) {
|
||||
s := c.take()
|
||||
s.verbatim = true
|
||||
s.passthrough = true
|
||||
s.rawPkt = pkt
|
||||
c.slots = append(c.slots, s)
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) seed(pkt []byte, info *parsedTCP) {
|
||||
if info.hdrLen+info.payLen > tcpCoalesceBufSize {
|
||||
// Pathological shape that can't ride a superpacket; emit as-is. No chain for this flow can
|
||||
// be open here (commitParsed evicts before seeding), so sealFlow is defense in depth
|
||||
// against a stale cache entry absorbing later data.
|
||||
c.sealFlow(info.fk)
|
||||
c.addVerbatim(pkt)
|
||||
func (c *TCPCoalescer) seed(pkt []byte, info parsedTCP) {
|
||||
if info.hdrLen > tcpCoalesceHdrCap || info.hdrLen+info.payLen > tcpCoalesceBufSize {
|
||||
// Pathological shape — can't fit our scratch, emit as-is.
|
||||
c.addPassthrough(pkt)
|
||||
return
|
||||
}
|
||||
s := c.take()
|
||||
s.verbatim = false
|
||||
// rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is
|
||||
// the superpacket header flushSlot patches in place.
|
||||
s.rawPkt = pkt
|
||||
s.passthrough = false
|
||||
s.rawPkt = nil
|
||||
copy(s.hdrBuf[:], pkt[:info.hdrLen])
|
||||
s.hdrLen = info.hdrLen
|
||||
s.ipHdrLen = info.ipHdrLen
|
||||
s.isV6 = info.fk.isV6
|
||||
@@ -282,23 +299,26 @@ func (c *TCPCoalescer) seed(pkt []byte, info *parsedTCP) {
|
||||
s.numSeg = 1
|
||||
s.totalPay = info.payLen
|
||||
s.nextSeq = info.seq + uint32(info.payLen)
|
||||
s.psh = info.flags&tcpFlagPsh != 0
|
||||
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||
c.slots = append(c.slots, s)
|
||||
if info.flags&tcpFlagPsh == 0 {
|
||||
if !s.psh {
|
||||
c.openSlots[info.fk] = s
|
||||
c.lastSlot = s
|
||||
} else {
|
||||
// PSH on the seed closes the chain immediately; it is never registered as open.
|
||||
// Drop any stale entry for this flow too (defense in depth, unreachable if lastSlot's lockstep invariant holds).
|
||||
c.sealFlow(info.fk)
|
||||
} else if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||
// PSH-on-seed seals the slot immediately. Any prior cached open
|
||||
// slot for this flow has just been sealed-and-replaced by this
|
||||
// passthrough-shaped seed, so drop the cache too.
|
||||
c.lastSlot = nil
|
||||
}
|
||||
}
|
||||
|
||||
// canAppend reports whether info's packet extends the slot's seed: same header shape and stable
|
||||
// contents, adjacent seq, not oversized. A closed chain never reaches here; closing removes the
|
||||
// slot from openSlots, the only path in. The header fields read from rawPkt are always pristine:
|
||||
// the only pre-flush mutation is the PSH propagate, which also closes the chain.
|
||||
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info *parsedTCP) bool {
|
||||
// canAppend reports whether info's packet extends the slot's seed: same
|
||||
// header shape and stable contents, adjacent seq, not oversized, chain not closed.
|
||||
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info parsedTCP) bool {
|
||||
if s.psh {
|
||||
return false
|
||||
}
|
||||
if info.hdrLen != s.hdrLen {
|
||||
return false
|
||||
}
|
||||
@@ -314,35 +334,31 @@ func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info *parsedTCP) b
|
||||
if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize {
|
||||
return false
|
||||
}
|
||||
// ECE state must be stable across a burst.
|
||||
// Receivers expect the flag set on every segment of a CE-echoing window or none.
|
||||
seedFlags := s.rawPkt[s.ipHdrLen+13]
|
||||
// ECE state must be stable across a burst — receivers expect the
|
||||
// flag set on every segment of a CE-echoing window or none.
|
||||
seedFlags := s.hdrBuf[s.ipHdrLen+13]
|
||||
if (seedFlags^info.flags)&tcpFlagEce != 0 {
|
||||
return false
|
||||
}
|
||||
if !s.isV6 && !ipv4CanCoalesceID(s.rawPkt, pkt, s.numSeg) {
|
||||
return false
|
||||
}
|
||||
if !headersMatch(s.rawPkt[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||
if !headersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// appendPayload folds info's packet into s and reports whether the chain is now closed: the
|
||||
// segment was sub-gsoSize (kernel TSO allows only the final segment to be short) or carried PSH.
|
||||
// The caller must deregister a closed slot from openSlots.
|
||||
func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info *parsedTCP) bool {
|
||||
func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info parsedTCP) {
|
||||
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||
s.numSeg++
|
||||
s.totalPay += info.payLen
|
||||
s.nextSeq = info.seq + uint32(info.payLen)
|
||||
if info.flags&tcpFlagPsh != 0 {
|
||||
// Propagate PSH into the seed header so kernel TSO sets it on the last segment. Mutating
|
||||
// rawPkt is safe: PSH also closes the chain, so no admission check re-reads this header.
|
||||
s.rawPkt[s.ipHdrLen+13] |= tcpFlagPsh
|
||||
// Propagate PSH into the seed header so kernel TSO sets it on the
|
||||
// last segment. Without this the sender's push signal is dropped.
|
||||
s.hdrBuf[s.ipHdrLen+13] |= tcpFlagPsh
|
||||
}
|
||||
if info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0 {
|
||||
s.psh = true
|
||||
}
|
||||
return info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) take() *coalesceSlot {
|
||||
@@ -356,18 +372,21 @@ func (c *TCPCoalescer) take() *coalesceSlot {
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) release(s *coalesceSlot) {
|
||||
s.passthrough = false
|
||||
s.rawPkt = nil
|
||||
clear(s.payIovs)
|
||||
*s = coalesceSlot{payIovs: s.payIovs[:0]}
|
||||
s.payIovs = s.payIovs[:0]
|
||||
s.numSeg = 0
|
||||
s.totalPay = 0
|
||||
s.psh = false
|
||||
c.pool = append(c.pool, s)
|
||||
}
|
||||
|
||||
// flushSlot patches the superpacket header in place in rawPkt (total length, IPv4 header
|
||||
// checksum, pseudo-header checksum seed) and calls WriteGSO. The slot is released right after,
|
||||
// so nothing re-reads the patched header. Does not remove the slot from c.slots.
|
||||
// flushSlot patches the header and calls WriteGSO. Does not remove the slot from c.slots.
|
||||
func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
|
||||
total := s.hdrLen + s.totalPay
|
||||
l4Len := total - s.ipHdrLen
|
||||
hdr := s.rawPkt[:s.hdrLen]
|
||||
hdr := s.hdrBuf[:s.hdrLen]
|
||||
|
||||
if s.isV6 {
|
||||
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
|
||||
@@ -387,7 +406,7 @@ func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
|
||||
tcsum := s.ipHdrLen + 16
|
||||
binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum))
|
||||
|
||||
return c.w.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoTCP)
|
||||
return c.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoTCP)
|
||||
}
|
||||
|
||||
// headersMatch compares two IP+TCP header prefixes for byte-for-byte
|
||||
@@ -418,6 +437,242 @@ func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// reorderForFlush neutralizes wire-side reorder that the rxOrder buffer
|
||||
// couldn't catch (anything crossing a recvmmsg batch boundary). Without
|
||||
// this pass a small wire reorder — counter 250 arriving in batch K when
|
||||
// 200..249 are coming in batch K+1 — would seed an out-of-seq slot first
|
||||
// and emit it ahead of the lower-seq slot, manifesting at the inner TCP
|
||||
// receiver as a much larger reorder than the wire actually had.
|
||||
//
|
||||
// Two phases:
|
||||
// 1. Sort each passthrough-bounded segment of c.slots by (flow, seq).
|
||||
// Cross-flow ordering inside a segment isn't preserved (it never was
|
||||
// and doesn't matter for any single flow's TCP correctness).
|
||||
// 2. Sweep once and merge adjacent same-flow slots whose ranges are now
|
||||
// contiguous AND whose tail is gsoSize-aligned. The tail constraint
|
||||
// matters because the kernel TSO splitter chops at gsoSize from the
|
||||
// start of the merged payload — a short segment in the middle would
|
||||
// desynchronize every later segment.
|
||||
//
|
||||
// Passthrough slots act as barriers: the merge check skips them on either
|
||||
// side, so a SYN/FIN/RST/CWR is never reordered relative to its flow's
|
||||
// data.
|
||||
func (c *TCPCoalescer) reorderForFlush() {
|
||||
if len(c.slots) <= 1 {
|
||||
return
|
||||
}
|
||||
runStart := 0
|
||||
for i := 0; i <= len(c.slots); i++ {
|
||||
if i < len(c.slots) && !c.slots[i].passthrough {
|
||||
continue
|
||||
}
|
||||
c.sortRun(c.slots[runStart:i])
|
||||
runStart = i + 1
|
||||
}
|
||||
out := c.slots[:0]
|
||||
logged := false
|
||||
for _, s := range c.slots {
|
||||
if n := len(out); n > 0 {
|
||||
prev := out[n-1]
|
||||
if !prev.passthrough && !s.passthrough && prev.fk == s.fk {
|
||||
// Same-flow neighbors after sort. If they aren't seq-
|
||||
// contiguous it's a real gap — packets the wire reordered
|
||||
// across batches, or actual loss before nebula. Log it so
|
||||
// the operator can quantify how often it happens; the data
|
||||
// itself still emits in seq order, kernel TCP handles the
|
||||
// gap via its OOO queue.
|
||||
if c.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
if prev.nextSeq != slotSeedSeq(s) {
|
||||
logged = true
|
||||
gap := int64(slotSeedSeq(s)) - int64(prev.nextSeq)
|
||||
c.l.Debug("tcp coalesce: cross-slot seq gap",
|
||||
"src", flowKeyAddr(s.fk, false),
|
||||
"dst", flowKeyAddr(s.fk, true),
|
||||
"sport", s.fk.sport,
|
||||
"dport", s.fk.dport,
|
||||
"prev_seed_seq", slotSeedSeq(prev),
|
||||
"prev_next_seq", prev.nextSeq,
|
||||
"this_seed_seq", slotSeedSeq(s),
|
||||
"gap_bytes", gap,
|
||||
"prev_seg_count", prev.numSeg,
|
||||
"prev_total_pay", prev.totalPay,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
if canMergeSlots(prev, s) {
|
||||
mergeSlots(prev, s)
|
||||
c.release(s)
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
out = append(out, s)
|
||||
}
|
||||
if logged {
|
||||
c.l.Warn("==== end of batch ====")
|
||||
}
|
||||
c.slots = out
|
||||
}
|
||||
|
||||
// flowKeyAddr returns the src or dst address from fk as a netip.Addr for
|
||||
// logging. Only used on the cold gap-log path so the netip allocation
|
||||
// doesn't matter.
|
||||
func flowKeyAddr(fk flowKey, dst bool) netip.Addr {
|
||||
src := fk.src
|
||||
if dst {
|
||||
src = fk.dst
|
||||
}
|
||||
if fk.isV6 {
|
||||
return netip.AddrFrom16(src)
|
||||
}
|
||||
var v4 [4]byte
|
||||
copy(v4[:], src[:4])
|
||||
return netip.AddrFrom4(v4)
|
||||
}
|
||||
|
||||
// sortRun stable-sorts run by (flowKey, seedSeq) so each flow's slots
|
||||
// cluster together in seq order, ready for the merge sweep. Stable so
|
||||
// equal-key slots keep their original relative position (defensive — a
|
||||
// duplicate seedSeq would already mean something's wrong upstream).
|
||||
func (c *TCPCoalescer) sortRun(run []*coalesceSlot) {
|
||||
if len(run) <= 1 {
|
||||
return
|
||||
}
|
||||
// slices.SortStableFunc with a free, non-capturing comparator avoids the
|
||||
// reflection + closure-escape allocations that sort.SliceStable forces.
|
||||
slices.SortStableFunc(run, compareCoalesceSlots)
|
||||
}
|
||||
|
||||
func compareCoalesceSlots(a, b *coalesceSlot) int {
|
||||
if cmp := flowKeyCompare(a.fk, b.fk); cmp != 0 {
|
||||
return cmp
|
||||
}
|
||||
aSeq, bSeq := slotSeedSeq(a), slotSeedSeq(b)
|
||||
if aSeq == bSeq {
|
||||
return 0
|
||||
}
|
||||
if tcpSeqLess(aSeq, bSeq) {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
}
|
||||
|
||||
// slotSeedSeq returns the TCP seq of the slot's seed (first segment).
|
||||
// nextSeq tracks the seq just past the last appended byte; subtracting
|
||||
// totalPay walks back to the seed. uint32 wraparound is the right TCP
|
||||
// arithmetic so no special-casing is needed.
|
||||
func slotSeedSeq(s *coalesceSlot) uint32 {
|
||||
return s.nextSeq - uint32(s.totalPay)
|
||||
}
|
||||
|
||||
// tcpSeqLess reports whether a precedes b in TCP serial-number arithmetic
|
||||
// (RFC 1323 §2.3). The signed int32 cast turns the modular subtraction
|
||||
// into the right comparison even across the 2^32 wrap.
|
||||
func tcpSeqLess(a, b uint32) bool {
|
||||
return int32(a-b) < 0
|
||||
}
|
||||
|
||||
// flowKeyCompare orders flowKeys deterministically. The exact ordering
|
||||
// is irrelevant — only that same-flow slots cluster together so the
|
||||
// post-sort sweep can merge contiguous pairs.
|
||||
func flowKeyCompare(a, b flowKey) int {
|
||||
// Cheap scalar fields first so most non-matching keys short-circuit
|
||||
// without ever calling bytes.Compare. sport is the ephemeral port on
|
||||
// egress flows and discriminates fastest. For matching keys (same
|
||||
// flow), array equality on src/dst inlines to word-sized compares,
|
||||
// so we only pay bytes.Compare when the arrays actually differ.
|
||||
if a.sport != b.sport {
|
||||
if a.sport < b.sport {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
}
|
||||
if a.dport != b.dport {
|
||||
if a.dport < b.dport {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
}
|
||||
if a.dst != b.dst {
|
||||
return bytes.Compare(a.dst[:], b.dst[:])
|
||||
}
|
||||
if a.src != b.src {
|
||||
return bytes.Compare(a.src[:], b.src[:])
|
||||
}
|
||||
if a.isV6 != b.isV6 {
|
||||
if !a.isV6 {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// canMergeSlots reports whether s can fold into prev as one merged TSO
|
||||
// superpacket. Same flow, contiguous TCP byte range, equal gsoSize, and
|
||||
// fits within the kernel TSO limits. The tail-of-prev check rejects any
|
||||
// merge whose first slot ended on a sub-gsoSize segment — kernel TSO
|
||||
// would split the merged skb at gsoSize boundaries from the start, so a
|
||||
// short segment in the middle would corrupt every later segment. PSH and
|
||||
// ECE state must agree across both slots: PSH is a semantic delimiter
|
||||
// (preserving the sender's push boundary) and ECE state must be uniform
|
||||
// across a window (the same rule canAppend enforces for in-flow appends).
|
||||
// The IP-level ECN codepoint must also match: this check calls headersMatch
|
||||
// → ipHeadersMatch, which compares the full DSCP/ECN byte, so two slots with
|
||||
// differing ECN marks stay separate superpackets, each keeping its own mark.
|
||||
//
|
||||
// Note: a slot sealed by reorder (canAppend returned false on seq
|
||||
// mismatch) keeps psh=false, so this restriction does not block the
|
||||
// reorder-fix merge — only legitimate PSH-set seals.
|
||||
func canMergeSlots(prev, s *coalesceSlot) bool {
|
||||
if prev.psh {
|
||||
return false
|
||||
}
|
||||
if prev.fk != s.fk {
|
||||
return false
|
||||
}
|
||||
if prev.gsoSize != s.gsoSize {
|
||||
return false
|
||||
}
|
||||
if prev.nextSeq != slotSeedSeq(s) {
|
||||
return false
|
||||
}
|
||||
if prev.numSeg+s.numSeg > tcpCoalesceMaxSegs {
|
||||
return false
|
||||
}
|
||||
if prev.hdrLen+prev.totalPay+s.totalPay > tcpCoalesceBufSize {
|
||||
return false
|
||||
}
|
||||
if len(prev.payIovs[len(prev.payIovs)-1]) != prev.gsoSize {
|
||||
return false
|
||||
}
|
||||
prevFlags := prev.hdrBuf[prev.ipHdrLen+13]
|
||||
sFlags := s.hdrBuf[s.ipHdrLen+13]
|
||||
if (prevFlags^sFlags)&tcpFlagEce != 0 {
|
||||
return false
|
||||
}
|
||||
if !headersMatch(prev.hdrBuf[:prev.hdrLen], s.hdrBuf[:s.hdrLen], prev.isV6, prev.ipHdrLen) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// mergeSlots folds src into dst in place: payIovs concatenated, counters
|
||||
// and totals updated, PSH OR'd into the seed header so the push signal is
|
||||
// not lost. The seed header's seq, gsoSize, and fk are unchanged. Caller
|
||||
// is responsible for releasing src (it's no longer in c.slots after this call).
|
||||
func mergeSlots(dst, src *coalesceSlot) {
|
||||
dst.payIovs = append(dst.payIovs, src.payIovs...)
|
||||
dst.numSeg += src.numSeg
|
||||
dst.totalPay += src.totalPay
|
||||
dst.nextSeq = src.nextSeq
|
||||
if src.psh {
|
||||
dst.psh = true
|
||||
dst.hdrBuf[dst.ipHdrLen+13] |= tcpFlagPsh
|
||||
}
|
||||
}
|
||||
|
||||
// ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must
|
||||
// already have its checksum field zeroed) and returns the folded/inverted
|
||||
// 16-bit value to store.
|
||||
@@ -462,8 +717,9 @@ func pseudoSumIPv6(src, dst []byte, proto byte, l4Len int) uint32 {
|
||||
return sum
|
||||
}
|
||||
|
||||
// foldOnceNoInvert folds the 32-bit accumulator to 16 bits and returns it unchanged (no one's complement).
|
||||
// This is what virtio NEEDS_CSUM wants in the L4 checksum field
|
||||
// foldOnceNoInvert folds the 32-bit accumulator to 16 bits and returns it
|
||||
// unchanged (no one's complement). This is what virtio NEEDS_CSUM wants in
|
||||
// the L4 checksum field — the kernel will add the payload sum and invert.
|
||||
func foldOnceNoInvert(sum uint32) uint16 {
|
||||
for sum>>16 != 0 {
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
|
||||
@@ -2,9 +2,9 @@ package batch
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"runtime"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/test"
|
||||
)
|
||||
@@ -55,31 +55,7 @@ func buildTCPv4Interleaved(nFlows, perFlow, payloadLen int) [][]byte {
|
||||
return pkts
|
||||
}
|
||||
|
||||
// buildTCPv4RunInterleaved returns nFlows*perFlow packets delivered in
|
||||
// runs of runLen per flow — the arrival pattern wire-side GRO actually
|
||||
// produces (deliverSegments splits each superdatagram into up to 64
|
||||
// same-flow packets back to back). Contrast with buildTCPv4Interleaved's
|
||||
// per-packet round-robin, the adversarial worst case for a last-slot cache.
|
||||
func buildTCPv4RunInterleaved(nFlows, perFlow, runLen, payloadLen int) [][]byte {
|
||||
pay := make([]byte, payloadLen)
|
||||
seqs := make([]uint32, nFlows)
|
||||
for i := range seqs {
|
||||
seqs[i] = uint32(1000 + i*1000000)
|
||||
}
|
||||
pkts := make([][]byte, 0, nFlows*perFlow)
|
||||
for done := 0; done < perFlow; done += runLen {
|
||||
for f := range nFlows {
|
||||
sport := uint16(10000 + f)
|
||||
for range runLen {
|
||||
pkts = append(pkts, buildTCPv4Ports(sport, 2000, seqs[f], tcpAck, pay))
|
||||
seqs[f] += uint32(payloadLen)
|
||||
}
|
||||
}
|
||||
}
|
||||
return pkts
|
||||
}
|
||||
|
||||
// buildICMPv4 returns a minimal non-TCP packet that takes the verbatim
|
||||
// buildICMPv4 returns a minimal non-TCP packet that takes the passthrough
|
||||
// branch in Commit.
|
||||
func buildICMPv4() []byte {
|
||||
pkt := make([]byte, 28)
|
||||
@@ -95,7 +71,8 @@ func buildICMPv4() []byte {
|
||||
// between batches, and reports per-packet cost.
|
||||
func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||
b.Helper()
|
||||
c := newTestTCPCoalescer(b, nopTunWriter{})
|
||||
arena := NewArena(0)
|
||||
c := NewTCPCoalescer(nopTunWriter{}, test.NewLogger(), arena.Reserve, arena.Reset)
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkts[0])))
|
||||
b.ResetTimer()
|
||||
@@ -136,17 +113,8 @@ func BenchmarkCommitInterleaved16(b *testing.B) {
|
||||
runCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkCommitRunInterleaved4 is 4 concurrent flows arriving in
|
||||
// GRO-burst runs of 16 — the realistic multi-flow pattern. A last-slot
|
||||
// cache hits for the length of each run; the per-packet round-robin
|
||||
// benches above are its worst case.
|
||||
func BenchmarkCommitRunInterleaved4(b *testing.B) {
|
||||
pkts := buildTCPv4RunInterleaved(4, tcpCoalesceMaxSegs, 16, 1200)
|
||||
runCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkCommitPassthrough exercises the non-TCP branch: parseBase
|
||||
// bails early and addVerbatim is the only work.
|
||||
// BenchmarkCommitPassthrough exercises the non-TCP branch: parseTCPBase
|
||||
// bails early and addPassthrough is the only work.
|
||||
func BenchmarkCommitPassthrough(b *testing.B) {
|
||||
pkt := buildICMPv4()
|
||||
pkts := make([][]byte, 64)
|
||||
@@ -158,7 +126,7 @@ func BenchmarkCommitPassthrough(b *testing.B) {
|
||||
|
||||
// BenchmarkCommitNonCoalesceableTCP sends SYN|ACK packets on one flow.
|
||||
// Each packet takes the "TCP but not admissible" branch which does a
|
||||
// map delete + verbatim. Measures the seal-without-slot cost.
|
||||
// map delete + passthrough. Measures the seal-without-slot cost.
|
||||
func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
|
||||
pay := make([]byte, 0)
|
||||
pkts := make([][]byte, 64)
|
||||
@@ -168,24 +136,18 @@ func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
|
||||
runCommitBench(b, pkts, 64)
|
||||
}
|
||||
|
||||
// runMultiCommitBench drives MultiCoalescer.Commit with in-order keys, so
|
||||
// it includes the staging sort's already-sorted fast path plus the
|
||||
// dispatch-time parse — the full steady-state cost of the batcher. The
|
||||
// ParsedPackets are precomputed: in production they fall out of the
|
||||
// firewall's newPacket, which this bench does not model.
|
||||
// runMultiCommitBench drives MultiCoalescer.Commit. The dispatcher does
|
||||
// the IP/L4 parse once and passes the parsed struct to the lane, so this
|
||||
// is the bench that shows the savings of skipping the lane's re-parse.
|
||||
func runMultiCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||
b.Helper()
|
||||
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger())
|
||||
pps := make([]*firewall.ParsedPacket, len(pkts))
|
||||
for i, p := range pkts {
|
||||
pps[i] = testPP(p)
|
||||
}
|
||||
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger(), NewArena(0), true, true)
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkts[0])))
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
j := i % len(pkts)
|
||||
if err := m.Commit(pkts[j], SortKey{Epoch: 1, Counter: uint64(i + 1)}, pps[j]); err != nil {
|
||||
pkt := pkts[i%len(pkts)]
|
||||
if err := m.Commit(pkt); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if (i+1)%batchSize == 0 {
|
||||
@@ -212,3 +174,68 @@ func BenchmarkMultiCommitInterleaved4(b *testing.B) {
|
||||
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
|
||||
runMultiCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// flowKeyPair is one comparison input for the flowKeyCompare bench.
|
||||
type flowKeyPair struct{ a, b flowKey }
|
||||
|
||||
// makeFlowKey builds an IPv4 flowKey from compact inputs.
|
||||
func makeFlowKey(srcLow, dstLow uint32, sport, dport uint16) flowKey {
|
||||
var fk flowKey
|
||||
binary.BigEndian.PutUint32(fk.src[12:16], srcLow)
|
||||
binary.BigEndian.PutUint32(fk.dst[12:16], dstLow)
|
||||
fk.sport = sport
|
||||
fk.dport = dport
|
||||
return fk
|
||||
}
|
||||
|
||||
// flowKeyCases are the workload mixes flowKeyCompare sees in practice.
|
||||
// - sameFlow: equal keys; tests the equal-path cost (sort runs hit this
|
||||
// repeatedly when many segments share a flow).
|
||||
// - sportDiffers: same src/dst/dport, different sport — the typical
|
||||
// "sibling flows from one host to one server" pattern.
|
||||
// - dstDiffers: same src/sport/dport, different dst — outbound to many
|
||||
// servers from a fixed local port.
|
||||
// - allDiffer: every field differs; worst case for short-circuiting.
|
||||
func flowKeyCases() map[string][]flowKeyPair {
|
||||
const n = 64
|
||||
cases := map[string][]flowKeyPair{
|
||||
"sameFlow": make([]flowKeyPair, n),
|
||||
"sportDiffers": make([]flowKeyPair, n),
|
||||
"dstDiffers": make([]flowKeyPair, n),
|
||||
"allDiffer": make([]flowKeyPair, n),
|
||||
}
|
||||
for i := range n {
|
||||
base := makeFlowKey(0x0a000001, 0x0a000002, 40000, 443)
|
||||
cases["sameFlow"][i] = flowKeyPair{a: base, b: base}
|
||||
cases["sportDiffers"][i] = flowKeyPair{
|
||||
a: base,
|
||||
b: makeFlowKey(0x0a000001, 0x0a000002, uint16(40001+i), 443),
|
||||
}
|
||||
cases["dstDiffers"][i] = flowKeyPair{
|
||||
a: base,
|
||||
b: makeFlowKey(0x0a000001, uint32(0x0a000002+i+1), 40000, 443),
|
||||
}
|
||||
cases["allDiffer"][i] = flowKeyPair{
|
||||
a: makeFlowKey(uint32(0x0a000001+i), uint32(0x0a000002+i), uint16(40000+i), uint16(80+i)),
|
||||
b: makeFlowKey(uint32(0x0b000001+i), uint32(0x0b000002+i), uint16(50000+i), uint16(443+i)),
|
||||
}
|
||||
}
|
||||
return cases
|
||||
}
|
||||
|
||||
// BenchmarkFlowKeyCompare measures flowKeyCompare across the workloads
|
||||
// the sort step actually sees. Use this to compare reorderings.
|
||||
func BenchmarkFlowKeyCompare(b *testing.B) {
|
||||
for name, pairs := range flowKeyCases() {
|
||||
b.Run(name, func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
var sink int
|
||||
for i := 0; i < b.N; i++ {
|
||||
p := pairs[i&(len(pairs)-1)]
|
||||
sink += flowKeyCompare(p.a, p.b)
|
||||
}
|
||||
runtime.KeepAlive(sink)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
+330
-472
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.
|
||||
type batchWriter interface {
|
||||
WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error)
|
||||
WriteBatch(bufs [][]byte, addrs []netip.AddrPort, outerECNs []byte) error
|
||||
}
|
||||
|
||||
// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch.
|
||||
@@ -16,6 +16,7 @@ type SendBatch struct {
|
||||
out batchWriter
|
||||
bufs [][]byte
|
||||
dsts []netip.AddrPort
|
||||
ecns []byte
|
||||
arena *Arena
|
||||
}
|
||||
|
||||
@@ -25,6 +26,7 @@ func NewSendBatch(out batchWriter, batchCap, arenaSize int) *SendBatch {
|
||||
out: out,
|
||||
bufs: make([][]byte, 0, batchCap),
|
||||
dsts: make([]netip.AddrPort, 0, batchCap),
|
||||
ecns: make([]byte, 0, batchCap),
|
||||
arena: NewArena(arenaSize),
|
||||
}
|
||||
}
|
||||
@@ -38,22 +40,21 @@ func (b *SendBatch) Reserve(sz int) []byte {
|
||||
// bounding how long the first packet of a large read batch waits.
|
||||
func (b *SendBatch) Len() int { return len(b.bufs) }
|
||||
|
||||
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort) {
|
||||
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort, outerECN byte) {
|
||||
b.bufs = append(b.bufs, pkt)
|
||||
b.dsts = append(b.dsts, dst)
|
||||
b.ecns = append(b.ecns, outerECN)
|
||||
}
|
||||
|
||||
// Flush writes every queued packet and reports how many actually went out. A short count means some destinations
|
||||
// were undeliverable; the batch is drained either way.
|
||||
func (b *SendBatch) Flush() (int, error) {
|
||||
func (b *SendBatch) Flush() error {
|
||||
var err error
|
||||
written := 0
|
||||
if len(b.bufs) > 0 {
|
||||
written, err = b.out.WriteBatch(b.bufs, b.dsts)
|
||||
err = b.out.WriteBatch(b.bufs, b.dsts, b.ecns)
|
||||
}
|
||||
clear(b.bufs)
|
||||
b.bufs = b.bufs[:0]
|
||||
b.dsts = b.dsts[:0]
|
||||
b.ecns = b.ecns[:0]
|
||||
b.arena.Reset()
|
||||
return written, err
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -8,9 +8,10 @@ import (
|
||||
type fakeBatchWriter struct {
|
||||
bufs [][]byte
|
||||
addrs []netip.AddrPort
|
||||
ecns []byte
|
||||
}
|
||||
|
||||
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
|
||||
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, ecns []byte) error {
|
||||
// Snapshot — SendBatch.Flush nils its slot pointers right after WriteBatch
|
||||
// returns, so tests must capture data before that happens.
|
||||
w.bufs = make([][]byte, len(bufs))
|
||||
@@ -20,7 +21,8 @@ func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int
|
||||
w.bufs[i] = cp
|
||||
}
|
||||
w.addrs = append(w.addrs[:0], addrs...)
|
||||
return len(bufs), nil
|
||||
w.ecns = append(w.ecns[:0], ecns...)
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestSendBatchReserveCommitFlush(t *testing.T) {
|
||||
@@ -34,9 +36,9 @@ func TestSendBatchReserveCommitFlush(t *testing.T) {
|
||||
t.Fatalf("slot %d: cap=%d want 32", i, cap(slot))
|
||||
}
|
||||
pkt := append(slot[:0], byte(i), byte(i+1), byte(i+2))
|
||||
b.Commit(pkt, ap)
|
||||
b.Commit(pkt, ap, 0)
|
||||
}
|
||||
if _, err := b.Flush(); err != nil {
|
||||
if err := b.Flush(); err != nil {
|
||||
t.Fatalf("Flush: %v", err)
|
||||
}
|
||||
if len(fw.bufs) != 4 {
|
||||
@@ -53,7 +55,7 @@ func TestSendBatchReserveCommitFlush(t *testing.T) {
|
||||
|
||||
// Flush again with nothing committed — should be a no-op.
|
||||
fw.bufs = nil
|
||||
if _, err := b.Flush(); err != nil {
|
||||
if err := b.Flush(); err != nil {
|
||||
t.Fatalf("empty Flush: %v", err)
|
||||
}
|
||||
if fw.bufs != nil {
|
||||
@@ -75,9 +77,9 @@ func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
|
||||
for i := 0; i < 3; i++ {
|
||||
s := b.Reserve(8)
|
||||
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
|
||||
b.Commit(pkt, ap)
|
||||
b.Commit(pkt, ap, 0)
|
||||
}
|
||||
if _, err := b.Flush(); err != nil {
|
||||
if err := b.Flush(); err != nil {
|
||||
t.Fatalf("Flush: %v", err)
|
||||
}
|
||||
|
||||
@@ -96,18 +98,18 @@ func TestSendBatchGrowPreservesCommitted(t *testing.T) {
|
||||
|
||||
s1 := b.Reserve(4)
|
||||
pkt1 := append(s1[:0], 0x11, 0x22, 0x33, 0x44)
|
||||
b.Commit(pkt1, ap)
|
||||
b.Commit(pkt1, ap, 0)
|
||||
|
||||
s2 := b.Reserve(8) // exceeds remaining cap, triggers grow
|
||||
pkt2 := append(s2[:0], 0xA, 0xB, 0xC, 0xD, 0xE)
|
||||
b.Commit(pkt2, ap)
|
||||
b.Commit(pkt2, ap, 0)
|
||||
|
||||
// pkt1 must still be intact even though backing reallocated.
|
||||
if pkt1[0] != 0x11 || pkt1[3] != 0x44 {
|
||||
t.Fatalf("first packet corrupted by grow: %x", pkt1)
|
||||
}
|
||||
|
||||
if _, err := b.Flush(); err != nil {
|
||||
if err := b.Flush(); err != nil {
|
||||
t.Fatalf("Flush: %v", err)
|
||||
}
|
||||
if len(fw.bufs) != 2 {
|
||||
|
||||
+152
-142
@@ -1,7 +1,6 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
|
||||
@@ -19,56 +18,72 @@ const udpCoalesceBufSize = 65535
|
||||
// accepts up to 64 segments per skb (UDP_MAX_SEGMENTS); stay under that.
|
||||
const udpCoalesceMaxSegs = 64
|
||||
|
||||
// udpCoalesceHdrCap is the scratch space we copy a seed's IP+UDP header
|
||||
// into. IPv6 (40) + UDP (8) = 48; round up for safety.
|
||||
const udpCoalesceHdrCap = 64
|
||||
|
||||
// udpSlot is one entry in the UDPCoalescer's ordered event queue.
|
||||
type udpSlot struct {
|
||||
verbatim bool
|
||||
// rawPkt is borrowed: the whole packet for verbatim slots, the seed
|
||||
// packet for coalesce slots. A coalesce slot that never grows past one
|
||||
// segment is emitted from rawPkt so its original (already valid) L4
|
||||
// checksum ships DATA_VALID instead of making the kernel recompute it.
|
||||
// A multi-segment slot's superpacket header is rawPkt's, patched in place at flush.
|
||||
rawPkt []byte
|
||||
passthrough bool
|
||||
rawPkt []byte // borrowed when passthrough
|
||||
|
||||
fk flowKey
|
||||
hdrBuf [udpCoalesceHdrCap]byte
|
||||
hdrLen int
|
||||
ipHdrLen int
|
||||
isV6 bool
|
||||
gsoSize int // per-segment UDP payload length
|
||||
numSeg int
|
||||
totalPay int
|
||||
payIovs [][]byte
|
||||
// sealed closes the chain: set when a sub-gsoSize segment is appended
|
||||
// (kernel UDP-GSO requires every segment but the last to be exactly
|
||||
// gsoSize) or when limits are hit. No further appends after.
|
||||
sealed bool
|
||||
payIovs [][]byte
|
||||
}
|
||||
|
||||
// UDPCoalescer accumulates adjacent in-flow UDP datagrams across multiple
|
||||
// concurrent flows and emits each flow's run as a single GSO_UDP_L4 superpacket via tio.GSOWriter.
|
||||
// Preserves the in-flow order of packets as they are Commit-ed
|
||||
// concurrent flows and emits each flow's run as a single GSO_UDP_L4
|
||||
// superpacket via tio.GSOWriter. Falls back to per-packet writes when the
|
||||
// underlying writer doesn't support USO.
|
||||
//
|
||||
// All output — coalesced or not — is deferred until Flush so per-flow
|
||||
// arrival order is preserved on the wire. Cross-flow order is NOT preserved
|
||||
// across the TCP/UDP/passthrough split when this coalescer runs alongside
|
||||
// others — see multi_coalesce.go. Per-flow order is preserved because a
|
||||
// single 5-tuple only ever lands in one lane and each lane preserves its
|
||||
// own slot order.
|
||||
//
|
||||
// Owns no locks; one coalescer per TUN write queue.
|
||||
type UDPCoalescer struct {
|
||||
w tio.GSOWriter
|
||||
plainW io.Writer
|
||||
gsoW tio.GSOWriter // nil when the queue can't accept GSO_UDP_L4
|
||||
|
||||
slots []*udpSlot
|
||||
openSlots map[flowKey]*udpSlot
|
||||
// lastSlot caches the most recently touched open slot; see the
|
||||
// TCPCoalescer field of the same name. Single-flow QUIC bulk is the
|
||||
// dominant USO workload, and multi-flow arrival comes in GRO runs, so
|
||||
// the fk compare beats the map's 38-byte key hash on most packets.
|
||||
// Kept in lockstep with openSlots: nil whenever the slot it pointed at
|
||||
// is removed.
|
||||
lastSlot *udpSlot
|
||||
pool []*udpSlot
|
||||
pool []*udpSlot
|
||||
reserver Reserver
|
||||
resetter Resetter
|
||||
}
|
||||
|
||||
func NewUDPCoalescer(w io.Writer) *UDPCoalescer {
|
||||
gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
return &UDPCoalescer{
|
||||
w: gw,
|
||||
// NewUDPCoalescer wraps w. The caller is responsible for only constructing
|
||||
// this when the underlying Queue's Capabilities advertise USO; otherwise
|
||||
// the kernel may reject GSO_UDP_L4 writes. If w does not implement
|
||||
// tio.GSOWriter at all (single-packet Queue), the coalescer degrades to
|
||||
// plain Writes — same defensive shape as the TCP coalescer.
|
||||
func NewUDPCoalescer(w io.Writer, reserver Reserver, resetter Resetter) *UDPCoalescer {
|
||||
c := &UDPCoalescer{
|
||||
plainW: w,
|
||||
slots: make([]*udpSlot, 0, initialSlots),
|
||||
openSlots: make(map[flowKey]*udpSlot, initialSlots),
|
||||
pool: make([]*udpSlot, 0, initialSlots),
|
||||
reserver: reserver,
|
||||
resetter: resetter,
|
||||
}
|
||||
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP); ok {
|
||||
c.gsoW = gw
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
// parsedUDP holds the fields extracted from a single parse so later steps
|
||||
@@ -80,111 +95,104 @@ type parsedUDP struct {
|
||||
payLen int
|
||||
}
|
||||
|
||||
// parseAt extracts the flow key and IP/UDP offsets for a packet the dispatcher already knows is
|
||||
// UDP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on
|
||||
// entry and is filled in place. Returns false for malformed input or any shape that must not
|
||||
// coalesce (IPv4 options/fragmentation, IPv6 extension headers).
|
||||
func (p *parsedUDP) parseAt(pkt []byte, ipHdrLen int) bool {
|
||||
trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen)
|
||||
// parseUDP extracts the flow key and IP/UDP offsets for a UDP packet.
|
||||
// Returns ok=false for non-UDP, malformed, or unsupported header shapes
|
||||
// (IPv4 with options/fragmentation, IPv6 with extension headers).
|
||||
func parseUDP(pkt []byte) (parsedUDP, bool) {
|
||||
var p parsedUDP
|
||||
ip, ok := parseIPPrologue(pkt, ipProtoUDP)
|
||||
if !ok {
|
||||
return false
|
||||
return p, false
|
||||
}
|
||||
return p.parseTail(trimmed, ipHdrLen)
|
||||
}
|
||||
pkt = ip.pkt
|
||||
p.fk = ip.fk
|
||||
p.ipHdrLen = ip.ipHdrLen
|
||||
|
||||
// parseTail layers the UDP-header parse on a validated IP prologue. pkt is the trimmed packet;
|
||||
// fk's addresses are already filled.
|
||||
func (p *parsedUDP) parseTail(pkt []byte, ipHdrLen int) bool {
|
||||
if len(pkt) < ipHdrLen+8 {
|
||||
return false
|
||||
if len(pkt) < p.ipHdrLen+8 {
|
||||
return p, false
|
||||
}
|
||||
p.hdrLen = p.ipHdrLen + 8
|
||||
// UDP `length` field: must equal IP-derived length-of-UDP-header-plus-payload.
|
||||
udpLen := int(binary.BigEndian.Uint16(pkt[ipHdrLen+4 : ipHdrLen+6]))
|
||||
if udpLen < 8 || udpLen > len(pkt)-ipHdrLen {
|
||||
return false
|
||||
udpLen := int(binary.BigEndian.Uint16(pkt[p.ipHdrLen+4 : p.ipHdrLen+6]))
|
||||
if udpLen < 8 || udpLen > len(pkt)-p.ipHdrLen {
|
||||
return p, false
|
||||
}
|
||||
p.ipHdrLen = ipHdrLen
|
||||
p.hdrLen = ipHdrLen + 8
|
||||
p.payLen = udpLen - 8
|
||||
p.fk.sport = binary.BigEndian.Uint16(pkt[ipHdrLen : ipHdrLen+2])
|
||||
p.fk.dport = binary.BigEndian.Uint16(pkt[ipHdrLen+2 : ipHdrLen+4])
|
||||
return true
|
||||
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2])
|
||||
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
|
||||
return p, true
|
||||
}
|
||||
|
||||
// sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips
|
||||
// hashing the 38-byte key when no chains are open.
|
||||
func (c *UDPCoalescer) sealFlow(fk flowKey) {
|
||||
if len(c.openSlots) == 0 {
|
||||
return
|
||||
}
|
||||
if last := c.lastSlot; last != nil && last.fk == fk {
|
||||
c.lastSlot = nil
|
||||
}
|
||||
delete(c.openSlots, fk)
|
||||
func (c *UDPCoalescer) Reserve(sz int) []byte {
|
||||
return c.reserver(sz)
|
||||
}
|
||||
|
||||
// commitStaged commits one staged packet dispatch routed to this lane. A shape the lane cannot
|
||||
// coalesce (any fragmentation, unparseable header) seals every open chain — its flow is unknown —
|
||||
// and rides the lane as an in-lane verbatim, still in transmission order.
|
||||
func (c *UDPCoalescer) commitStaged(sp stagedPacket) error {
|
||||
if sp.fragAny {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(sp.pkt)
|
||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
||||
func (c *UDPCoalescer) Commit(pkt []byte) error {
|
||||
if c.gsoW == nil {
|
||||
c.addPassthrough(pkt)
|
||||
return nil
|
||||
}
|
||||
var info parsedUDP
|
||||
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(sp.pkt)
|
||||
info, ok := parseUDP(pkt)
|
||||
if !ok {
|
||||
c.addPassthrough(pkt)
|
||||
return nil
|
||||
}
|
||||
return c.commitParsed(sp.pkt, &info)
|
||||
return c.commitParsed(pkt, info)
|
||||
}
|
||||
|
||||
// commitParsed commits one parsed UDP packet. The caller (dispatch, via parseAt) supplies a
|
||||
// valid parse so the header is not re-walked here.
|
||||
func (c *UDPCoalescer) commitParsed(pkt []byte, info *parsedUDP) error {
|
||||
// A zero-length UDP datagram (length == 8) is legal and must reach the TUN, but cannot be
|
||||
// coalesced.
|
||||
// commitParsed is the post-parse half of Commit. The caller must have
|
||||
// already verified parseUDP succeeded. Used by MultiCoalescer.Commit to
|
||||
// avoid re-walking the IP/UDP header.
|
||||
func (c *UDPCoalescer) commitParsed(pkt []byte, info parsedUDP) error {
|
||||
if c.gsoW == nil {
|
||||
c.addPassthrough(pkt)
|
||||
return nil
|
||||
}
|
||||
// A zero-length UDP datagram (UDP `length` == 8) is legal and must still
|
||||
// reach the TUN, but it can't be coalesced: a GSO slot would store an
|
||||
// empty payload iovec and the kernel has nothing to segment. Seal any
|
||||
// open chain for this flow (so a later, non-empty datagram seeds fresh
|
||||
// *after* this one and per-flow arrival order is preserved) and deliver
|
||||
// it as a plain single datagram.
|
||||
if info.payLen == 0 {
|
||||
c.sealFlow(info.fk)
|
||||
c.addVerbatim(pkt)
|
||||
delete(c.openSlots, info.fk)
|
||||
c.addPassthrough(pkt)
|
||||
return nil
|
||||
}
|
||||
// Cached-slot fast path; see the TCPCoalescer equivalent.
|
||||
var open *udpSlot
|
||||
if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||
open = last
|
||||
} else {
|
||||
open = c.openSlots[info.fk]
|
||||
}
|
||||
if open != nil {
|
||||
if open := c.openSlots[info.fk]; open != nil {
|
||||
if c.canAppend(open, pkt, info) {
|
||||
if c.appendPayload(open, pkt, info) {
|
||||
// Chain closed (short segment): stop extending it.
|
||||
c.sealFlow(info.fk)
|
||||
} else {
|
||||
c.lastSlot = open
|
||||
c.appendPayload(open, pkt, info)
|
||||
if open.sealed {
|
||||
delete(c.openSlots, info.fk)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
// Can't extend: evict it from openSlots and fall through to seed a
|
||||
// fresh slot.
|
||||
c.sealFlow(info.fk)
|
||||
// Can't extend — seal it and fall through to seed a fresh slot.
|
||||
delete(c.openSlots, info.fk)
|
||||
}
|
||||
c.seed(pkt, info)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Flush drains every queued slot and calls the configured Resetter.
|
||||
func (c *UDPCoalescer) Flush() error {
|
||||
first := c.drain()
|
||||
if c.resetter != nil {
|
||||
c.resetter()
|
||||
}
|
||||
return first
|
||||
}
|
||||
|
||||
// drain emits every queued slot in arrival order and clears the slot state.
|
||||
// It does NOT reset the arena: borrowed payload slices stay valid until the
|
||||
// arena's owner recycles it.
|
||||
func (c *UDPCoalescer) drain() error {
|
||||
var first error
|
||||
for _, s := range c.slots {
|
||||
var err error
|
||||
if s.verbatim || s.numSeg == 1 {
|
||||
// A slot that never grew is byte-identical to the packet it was
|
||||
// seeded from; ship the original so its valid checksum rides the
|
||||
// DATA_VALID path instead of paying a kernel software csum.
|
||||
_, err = c.w.Write(s.rawPkt)
|
||||
if s.passthrough {
|
||||
_, err = c.plainW.Write(s.rawPkt)
|
||||
} else {
|
||||
err = c.flushSlot(s)
|
||||
}
|
||||
@@ -196,38 +204,25 @@ func (c *UDPCoalescer) Flush() error {
|
||||
clear(c.slots)
|
||||
c.slots = c.slots[:0]
|
||||
clear(c.openSlots)
|
||||
c.lastSlot = nil
|
||||
return first
|
||||
}
|
||||
|
||||
// sealAllOpen closes every open coalesce chain. Called for unparseable packets: the flow key is
|
||||
// unknown, so any open chain could otherwise absorb later data and emit it ahead of this packet.
|
||||
func (c *UDPCoalescer) sealAllOpen() {
|
||||
clear(c.openSlots)
|
||||
c.lastSlot = nil
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) addVerbatim(pkt []byte) {
|
||||
func (c *UDPCoalescer) addPassthrough(pkt []byte) {
|
||||
s := c.take()
|
||||
s.verbatim = true
|
||||
s.passthrough = true
|
||||
s.rawPkt = pkt
|
||||
c.slots = append(c.slots, s)
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) seed(pkt []byte, info *parsedUDP) {
|
||||
if info.hdrLen+info.payLen > udpCoalesceBufSize {
|
||||
// Pathological shape that can't ride a superpacket; emit as-is. No chain for this flow can
|
||||
// be open here (commitParsed evicts before seeding), so sealFlow is defense in depth
|
||||
// against a stale cache entry absorbing later data.
|
||||
c.sealFlow(info.fk)
|
||||
c.addVerbatim(pkt)
|
||||
func (c *UDPCoalescer) seed(pkt []byte, info parsedUDP) {
|
||||
if info.hdrLen > udpCoalesceHdrCap || info.hdrLen+info.payLen > udpCoalesceBufSize {
|
||||
c.addPassthrough(pkt)
|
||||
return
|
||||
}
|
||||
s := c.take()
|
||||
s.verbatim = false
|
||||
// rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is
|
||||
// the superpacket header flushSlot patches in place.
|
||||
s.rawPkt = pkt
|
||||
s.passthrough = false
|
||||
s.rawPkt = nil
|
||||
copy(s.hdrBuf[:], pkt[:info.hdrLen])
|
||||
s.hdrLen = info.hdrLen
|
||||
s.ipHdrLen = info.ipHdrLen
|
||||
s.isV6 = info.fk.isV6
|
||||
@@ -235,16 +230,19 @@ func (c *UDPCoalescer) seed(pkt []byte, info *parsedUDP) {
|
||||
s.gsoSize = info.payLen
|
||||
s.numSeg = 1
|
||||
s.totalPay = info.payLen
|
||||
s.sealed = false
|
||||
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||
c.slots = append(c.slots, s)
|
||||
c.openSlots[info.fk] = s
|
||||
c.lastSlot = s
|
||||
}
|
||||
|
||||
// canAppend reports whether info's packet extends the slot's seed.
|
||||
// Kernel UDP-GSO requires every segment except possibly the last to be
|
||||
// exactly gsoSize, and the last may be shorter (≤ gsoSize).
|
||||
func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info *parsedUDP) bool {
|
||||
func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info parsedUDP) bool {
|
||||
if s.sealed {
|
||||
return false
|
||||
}
|
||||
if info.hdrLen != s.hdrLen {
|
||||
return false
|
||||
}
|
||||
@@ -257,25 +255,20 @@ func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info *parsedUDP) bool {
|
||||
if s.hdrLen+s.totalPay+info.payLen > udpCoalesceBufSize {
|
||||
return false
|
||||
}
|
||||
// Header reads use rawPkt, which is never mutated before flush. A closed chain never reaches
|
||||
// here; closing removes the slot from openSlots, the only path in.
|
||||
if !s.isV6 && !ipv4CanCoalesceID(s.rawPkt, pkt, s.numSeg) {
|
||||
return false
|
||||
}
|
||||
if !udpHeadersMatch(s.rawPkt[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||
if !udpHeadersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// appendPayload folds info's packet into s and reports whether the chain is now closed: kernel
|
||||
// UDP-GSO requires every segment but the last to be exactly gsoSize, so a short segment must be
|
||||
// the final one. The caller must deregister a closed slot from openSlots.
|
||||
func (c *UDPCoalescer) appendPayload(s *udpSlot, pkt []byte, info *parsedUDP) bool {
|
||||
func (c *UDPCoalescer) appendPayload(s *udpSlot, pkt []byte, info parsedUDP) {
|
||||
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||
s.numSeg++
|
||||
s.totalPay += info.payLen
|
||||
return info.payLen < s.gsoSize
|
||||
if info.payLen < s.gsoSize {
|
||||
// Last-segment-can-be-shorter: this seals the chain.
|
||||
s.sealed = true
|
||||
}
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) take() *udpSlot {
|
||||
@@ -289,19 +282,30 @@ func (c *UDPCoalescer) take() *udpSlot {
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) release(s *udpSlot) {
|
||||
// Reset every field, identity ones included; see TCPCoalescer.release.
|
||||
s.passthrough = false
|
||||
s.rawPkt = nil
|
||||
clear(s.payIovs)
|
||||
*s = udpSlot{payIovs: s.payIovs[:0]}
|
||||
s.payIovs = s.payIovs[:0]
|
||||
s.numSeg = 0
|
||||
s.totalPay = 0
|
||||
s.sealed = false
|
||||
c.pool = append(c.pool, s)
|
||||
}
|
||||
|
||||
// flushSlot patches the IP header total length / IPv6 payload length and
|
||||
// the UDP length to the *total* across all coalesced segments, then seeds
|
||||
// the UDP checksum field with the pseudo-header partial (single-fold, not
|
||||
// inverted) per virtio NEEDS_CSUM. The patches land in place in rawPkt; the
|
||||
// slot is released right after, so nothing re-reads the patched header.
|
||||
// inverted) per virtio NEEDS_CSUM. The kernel's ip_rcv_core (v4) and
|
||||
// ip6_rcv_core (v6) trim the skb to those length fields, so per-segment
|
||||
// values would silently drop everything but the first segment. The kernel
|
||||
// then walks each segment in __udp_gso_segment, recomputing per-segment
|
||||
// uh->len / iph->tot_len / IPv6 plen and adjusting the checksum via
|
||||
// `check = csum16_add(csum16_sub(uh->check, uh->len), newlen)` — meaning
|
||||
// our seed's uh->check must be consistent with the seed's uh->len, which
|
||||
// is what passing the total to both pseudoSum and the UDP length field
|
||||
// guarantees.
|
||||
func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
|
||||
hdr := s.rawPkt[:s.hdrLen]
|
||||
hdr := s.hdrBuf[:s.hdrLen]
|
||||
total := s.hdrLen + s.totalPay // full IP+UDP+all_payloads bytes
|
||||
l4Len := total - s.ipHdrLen // total UDP (8 + sum of payloads)
|
||||
|
||||
@@ -326,11 +330,14 @@ func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
|
||||
udpCsumOff := s.ipHdrLen + 6
|
||||
binary.BigEndian.PutUint16(hdr[udpCsumOff:udpCsumOff+2], foldOnceNoInvert(psum))
|
||||
|
||||
return c.w.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoUDP)
|
||||
return c.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoUDP)
|
||||
}
|
||||
|
||||
// udpHeadersMatch compares two IP+UDP header prefixes for byte-equality on
|
||||
// every field that must be identical across coalesced segments
|
||||
// every field that must be identical across coalesced segments. Length
|
||||
// fields are masked out (flushSlot rewrites them), but the IP-level ECN
|
||||
// codepoint is compared (via ipHeadersMatch) so segments with differing ECN
|
||||
// don't coalesce, matching kernel GRO.
|
||||
func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
@@ -338,8 +345,11 @@ func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
||||
if !ipHeadersMatch(a, b, isV6) {
|
||||
return false
|
||||
}
|
||||
// UDP: compare sport+dport ([0:4]). Skip length [4:6] and checksum [6:8]:
|
||||
// UDP: compare sport+dport ([0:4]). Skip length [4:6] and checksum [6:8] —
|
||||
// length varies (we rewrite at flush) and the checksum will be redone.
|
||||
udp := ipHdrLen
|
||||
return bytes.Equal(a[udp:udp+4], b[udp:udp+4])
|
||||
if a[udp] != b[udp] || a[udp+1] != b[udp+1] || a[udp+2] != b[udp+2] || a[udp+3] != b[udp+3] {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
@@ -1,72 +0,0 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
// buildUDPv4BulkFlow returns n equal-size datagrams on one flow — the
|
||||
// steady state for single-flow QUIC bulk, the workload USO exists for.
|
||||
func buildUDPv4BulkFlow(n, payloadLen int) [][]byte {
|
||||
pay := make([]byte, payloadLen)
|
||||
pkts := make([][]byte, n)
|
||||
for i := range pkts {
|
||||
pkts[i] = buildUDPv4(40000, 443, pay)
|
||||
}
|
||||
return pkts
|
||||
}
|
||||
|
||||
// buildUDPv4RunInterleaved mirrors buildTCPv4RunInterleaved: nFlows*perFlow
|
||||
// datagrams arriving in GRO-burst runs of runLen per flow.
|
||||
func buildUDPv4RunInterleaved(nFlows, perFlow, runLen, payloadLen int) [][]byte {
|
||||
pay := make([]byte, payloadLen)
|
||||
pkts := make([][]byte, 0, nFlows*perFlow)
|
||||
for done := 0; done < perFlow; done += runLen {
|
||||
for f := range nFlows {
|
||||
sport := uint16(40000 + f)
|
||||
for range runLen {
|
||||
pkts = append(pkts, buildUDPv4(sport, 443, pay))
|
||||
}
|
||||
}
|
||||
}
|
||||
return pkts
|
||||
}
|
||||
|
||||
// runUDPCommitBench drives UDPCoalescer.Commit over pkts batchSize at a
|
||||
// time, flushing between batches, and reports per-packet cost.
|
||||
func runUDPCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||
b.Helper()
|
||||
c := newTestUDPCoalescer(b, nopTunWriter{})
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkts[0])))
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
pkt := pkts[i%len(pkts)]
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if (i+1)%batchSize == 0 {
|
||||
if err := c.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
_ = c.Flush()
|
||||
}
|
||||
|
||||
// BenchmarkUDPCommitSingleFlow is the single-flow bulk steady state.
|
||||
func BenchmarkUDPCommitSingleFlow(b *testing.B) {
|
||||
pkts := buildUDPv4BulkFlow(udpCoalesceMaxSegs, 1200)
|
||||
runUDPCommitBench(b, pkts, udpCoalesceMaxSegs)
|
||||
}
|
||||
|
||||
// BenchmarkUDPCommitInterleaved4 is the adversarial per-packet round-robin.
|
||||
func BenchmarkUDPCommitInterleaved4(b *testing.B) {
|
||||
pkts := buildUDPv4RunInterleaved(4, udpCoalesceMaxSegs, 1, 1200)
|
||||
runUDPCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkUDPCommitRunInterleaved4 is 4 flows in GRO-burst runs of 16.
|
||||
func BenchmarkUDPCommitRunInterleaved4(b *testing.B) {
|
||||
pkts := buildUDPv4RunInterleaved(4, udpCoalesceMaxSegs, 16, 1200)
|
||||
runUDPCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
@@ -1,9 +1,7 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"testing"
|
||||
)
|
||||
|
||||
@@ -60,31 +58,29 @@ func buildUDPv6(sport, dport uint16, payload []byte) []byte {
|
||||
return pkt
|
||||
}
|
||||
|
||||
// newTestUDPCoalescer builds a coalescer over w and fails the test if w can't
|
||||
// do USO. See newTestTCPCoalescer.
|
||||
func newTestUDPCoalescer(tb testing.TB, w io.Writer) *UDPCoalescer {
|
||||
tb.Helper()
|
||||
c := NewUDPCoalescer(w)
|
||||
if c == nil {
|
||||
tb.Fatal("NewUDPCoalescer: writer does not support USO")
|
||||
func TestUDPCoalescerPassthroughWhenGSOUnavailable(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: false}
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
pkt := buildUDPv4(1000, 53, make([]byte, 100))
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
// TestNewUDPCoalescerRefusesWhenGSOUnavailable mirrors the TCP precondition:
|
||||
// no USO, no coalescer.
|
||||
func TestNewUDPCoalescerRefusesWhenGSOUnavailable(t *testing.T) {
|
||||
if c := NewUDPCoalescer(&fakeTunWriter{gsoEnabled: false}); c != nil {
|
||||
t.Fatalf("want nil for a non-USO writer, got %v", c)
|
||||
if len(w.writes) != 0 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("no Add-time writes: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
if c := NewUDPCoalescer(&plainOnlyWriter{}); c != nil {
|
||||
t.Fatalf("want nil for a plain writer, got %v", c)
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("want single plain write, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
}
|
||||
|
||||
func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
// ICMP packet
|
||||
pkt := make([]byte, 28)
|
||||
pkt[0] = 0x45
|
||||
@@ -105,7 +101,8 @@ func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
|
||||
|
||||
func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
pkt := buildUDPv4(1000, 53, make([]byte, 800))
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -113,21 +110,17 @@ func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// A slot that never grew past one datagram flushes as a plain Write of
|
||||
// the original packet bytes: the original (already valid) checksum
|
||||
// ships via the DATA_VALID path, so the kernel does no csum work.
|
||||
// WriteGSO is reserved for slots that actually coalesced (>=2 segs).
|
||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||
// Single-segment flush goes through WriteGSO; the writer infers GSO_NONE
|
||||
// from len(pays)==1 and the kernel fills in the UDP csum (NEEDS_CSUM).
|
||||
if len(w.gsoWrites) != 1 || len(w.writes) != 0 {
|
||||
t.Fatalf("single-seg flush: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
if !bytes.Equal(w.writes[0], pkt) {
|
||||
t.Errorf("plain write not byte-identical to committed packet: got %d bytes want %d", len(w.writes[0]), len(pkt))
|
||||
}
|
||||
}
|
||||
|
||||
func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
pay := make([]byte, 1200)
|
||||
for i := 0; i < 3; i++ {
|
||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||
@@ -167,7 +160,8 @@ func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
|
||||
// Last segment may be shorter, sealing the chain.
|
||||
func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
full := make([]byte, 1200)
|
||||
tail := make([]byte, 600)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||
@@ -186,23 +180,22 @@ func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// The sealed 3-datagram chain is a real superpacket; the re-seed stays
|
||||
// single-segment and flushes as a plain write of the original packet.
|
||||
if len(w.gsoWrites) != 1 || len(w.writes) != 1 {
|
||||
t.Fatalf("want 1 gso (sealed) + 1 plain (new seed), got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 gso writes (sealed + new seed), got %d", len(w.gsoWrites))
|
||||
}
|
||||
if len(w.gsoWrites[0].pays) != 3 {
|
||||
t.Errorf("super: want 3 pays, got %d", len(w.gsoWrites[0].pays))
|
||||
t.Errorf("first super: want 3 pays, got %d", len(w.gsoWrites[0].pays))
|
||||
}
|
||||
if got, want := len(w.writes[0]), 20+8+1200; got != want {
|
||||
t.Errorf("re-seed plain write len=%d want %d", got, want)
|
||||
if len(w.gsoWrites[1].pays) != 1 {
|
||||
t.Errorf("second super: want 1 pay (re-seed), got %d", len(w.gsoWrites[1].pays))
|
||||
}
|
||||
}
|
||||
|
||||
// A larger-than-gsoSize packet cannot extend the slot — it reseeds.
|
||||
func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -212,21 +205,16 @@ func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Both seeds stay single-segment → two plain writes in arrival order.
|
||||
if len(w.writes) != 2 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("want 2 separate plain writes, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
for i, want := range []int{20 + 8 + 800, 20 + 8 + 1200} {
|
||||
if len(w.writes[i]) != want {
|
||||
t.Errorf("write %d len=%d want %d", i, len(w.writes[i]), want)
|
||||
}
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 separate seeds, got %d", len(w.gsoWrites))
|
||||
}
|
||||
}
|
||||
|
||||
// Different 5-tuples must not coalesce.
|
||||
func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
pay := make([]byte, 800)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -257,7 +245,8 @@ func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
|
||||
// Caps at udpCoalesceMaxSegs.
|
||||
func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
pay := make([]byte, 100)
|
||||
for i := 0; i < udpCoalesceMaxSegs+5; i++ {
|
||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||
@@ -282,12 +271,12 @@ func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
|
||||
|
||||
// Differing IP ECN codepoints must not coalesce: udpHeadersMatch compares
|
||||
// the full ToS byte (matching kernel GRO). A CE-marked datagram mid-run
|
||||
// seals the Not-ECT chain and reseeds; the trailing Not-ECT datagram
|
||||
// reseeds again. All three stay single-segment, so each ships as a plain
|
||||
// write of its original bytes, keeping its own codepoint.
|
||||
// seals the Not-ECT chain and seeds a fresh superpacket that keeps CE; the
|
||||
// trailing Not-ECT datagram seeds another.
|
||||
func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
pay := make([]byte, 800)
|
||||
pkt0 := buildUDPv4(1000, 53, pay) // ECN=00 (Not-ECT)
|
||||
pkt1 := buildUDPv4(1000, 53, pay)
|
||||
@@ -301,13 +290,16 @@ func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 3 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("want 3 separate plain writes (differing ECN), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
if len(w.gsoWrites) != 3 {
|
||||
t.Fatalf("want 3 separate seeds (differing ECN), got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
wantECN := []byte{0x00, 0x03, 0x00}
|
||||
for i, p := range w.writes {
|
||||
if got := p[1] & 0x03; got != wantECN[i] {
|
||||
t.Errorf("write %d ECN=%#x want %#x", i, got, wantECN[i])
|
||||
for i, g := range w.gsoWrites {
|
||||
if len(g.pays) != 1 {
|
||||
t.Errorf("gso %d pay count=%d want 1", i, len(g.pays))
|
||||
}
|
||||
if got := g.hdr[1] & 0x03; got != wantECN[i] {
|
||||
t.Errorf("gso %d ECN=%#x want %#x", i, got, wantECN[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -315,7 +307,8 @@ func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
|
||||
// IPv6 path: same flow, equal-sized → coalesced.
|
||||
func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
pay := make([]byte, 1200)
|
||||
for i := 0; i < 3; i++ {
|
||||
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
|
||||
@@ -351,7 +344,8 @@ func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
|
||||
// DSCP differences must reseed: udpHeadersMatch compares the full ToS byte.
|
||||
func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
pay := make([]byte, 800)
|
||||
pkt0 := buildUDPv4(1000, 53, pay)
|
||||
pkt1 := buildUDPv4(1000, 53, pay)
|
||||
@@ -365,16 +359,16 @@ func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Both seeds stay single-segment → two plain writes, no gso.
|
||||
if len(w.writes) != 2 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("want 2 separate plain writes (different DSCP), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 separate seeds (different DSCP), got %d", len(w.gsoWrites))
|
||||
}
|
||||
}
|
||||
|
||||
// Fragmented IPv4 must not be coalesced.
|
||||
func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
||||
binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
@@ -395,7 +389,8 @@ func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
|
||||
// reach the GSO path. Regression: must not panic and must be written.
|
||||
func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
pkt := buildUDPv4(1000, 53, nil) // UDP length 8, zero payload
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -411,10 +406,11 @@ func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// IPv6 zero-length UDP datagram: same verbatim contract as v4.
|
||||
// IPv6 zero-length UDP datagram: same passthrough contract as v4.
|
||||
func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
pkt := buildUDPv6(1000, 53, nil) // UDP length 8, zero payload
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -435,7 +431,8 @@ func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
|
||||
// wire — per-flow arrival order (full, empty, full) must be preserved.
|
||||
func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
full := make([]byte, 800)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -450,23 +447,17 @@ func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// The empty datagram sealed the first slot, so the trailing full packet
|
||||
// can't join it. All three emit as plain writes (the two full datagrams
|
||||
// stayed single-segment; the empty one is verbatim) in per-flow
|
||||
// arrival order: full, empty, full.
|
||||
if len(w.writes) != 3 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("want 3 plain writes, got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
for i, want := range []int{20 + 8 + 800, 20 + 8, 20 + 8 + 800} {
|
||||
if len(w.writes[i]) != want {
|
||||
t.Errorf("write %d len=%d want %d (order full, empty, full)", i, len(w.writes[i]), want)
|
||||
}
|
||||
// can't join it: two single-segment superpackets bracket one plain write.
|
||||
if len(w.gsoWrites) != 2 || len(w.writes) != 1 {
|
||||
t.Fatalf("want 2 gso writes + 1 plain, got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
}
|
||||
|
||||
// IPv4 with options is not admissible (we require IHL=5).
|
||||
func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
arena := NewArena(0)
|
||||
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
||||
pkt[0] = 0x46 // IHL = 6 (24-byte IPv4 header — has options)
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
@@ -479,58 +470,3 @@ func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
|
||||
t.Fatalf("ipv4-with-options must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
}
|
||||
|
||||
// TestUDPCoalescerNonAtomicSequentialIDsCoalesce mirrors the TCP rule: DF
|
||||
// clear is fine as long as the IDs already run seed+1 per datagram, so
|
||||
// kernel USO's re-stamp reproduces them.
|
||||
func TestUDPCoalescerNonAtomicSequentialIDsCoalesce(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
|
||||
for i := range 2 {
|
||||
pkt := buildUDPv4(40000, 443, pay)
|
||||
setIPv4ID(pkt, uint16(40+i), false)
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 1 || len(w.gsoWrites[0].pays) != 2 {
|
||||
t.Fatalf("sequential-ID DF=0 datagrams must coalesce: gso=%d", len(w.gsoWrites))
|
||||
}
|
||||
}
|
||||
|
||||
// TestUDPCoalescerNonAtomicIDGapReseeds: an ID jump on a DF=0 flow breaks
|
||||
// the chain; each datagram stays a single-segment slot and flushes as a
|
||||
// plain write that keeps its own (meaningful) ID.
|
||||
func TestUDPCoalescerNonAtomicIDGapReseeds(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
|
||||
p1 := buildUDPv4(40000, 443, pay)
|
||||
setIPv4ID(p1, 40, false)
|
||||
p2 := buildUDPv4(40000, 443, pay)
|
||||
setIPv4ID(p2, 50, false)
|
||||
|
||||
if err := c.Commit(p1); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(p2); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 2 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("ID gap on DF=0 must reseed: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
for i, want := range []uint16{40, 50} {
|
||||
if id := binary.BigEndian.Uint16(w.writes[i][4:6]); id != want {
|
||||
t.Errorf("write %d: ID=%d want %d (must be preserved)", i, id, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -8,69 +8,37 @@ import (
|
||||
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||
)
|
||||
|
||||
// archImpl names one checksum function under test. The per-arch
|
||||
// export_*_test.go files enumerate the hand-written implementations so the
|
||||
// suite compares each one against gvisor directly, regardless of which one
|
||||
// the public Checksum dispatches to on the running CPU. Testing only the
|
||||
// dispatcher was tautological wherever it resolved to the gvisor fallback
|
||||
// (non-AVX2 amd64, fallback architectures) — gvisor compared with itself,
|
||||
// assembly untested, suite green.
|
||||
type archImpl struct {
|
||||
name string
|
||||
fn func([]byte, uint16) uint16
|
||||
available bool
|
||||
}
|
||||
|
||||
// implsUnderTest is the public dispatcher plus every arch implementation.
|
||||
func implsUnderTest() []archImpl {
|
||||
return append([]archImpl{{name: "dispatch", fn: Checksum, available: true}}, archImpls...)
|
||||
}
|
||||
|
||||
// requireAvailable skips loudly when the running CPU can't execute an
|
||||
// implementation — visible in test output, unlike the old silent tautology.
|
||||
func requireAvailable(t *testing.T, impl archImpl) {
|
||||
t.Helper()
|
||||
if !impl.available {
|
||||
t.Skipf("%s not supported on this CPU; its assembly is NOT tested in this run", impl.name)
|
||||
}
|
||||
}
|
||||
|
||||
// TestChecksumMatchesGvisor walks lengths from 0 to 4096, with several initial
|
||||
// seeds and a handful of starting alignments, asserting that each local
|
||||
// implementation matches gvisor's reference bit-for-bit.
|
||||
// seeds and a handful of starting alignments, asserting that our local
|
||||
// Checksum matches gvisor's reference bit-for-bit.
|
||||
func TestChecksumMatchesGvisor(t *testing.T) {
|
||||
for _, impl := range implsUnderTest() {
|
||||
t.Run(impl.name, func(t *testing.T) {
|
||||
requireAvailable(t, impl)
|
||||
rng := rand.New(rand.NewPCG(1, 2))
|
||||
const padFront = 16
|
||||
rng := rand.New(rand.NewPCG(1, 2))
|
||||
const padFront = 16
|
||||
|
||||
// Random pool large enough for the longest case + alignment slop.
|
||||
pool := make([]byte, 4096+padFront)
|
||||
for i := range pool {
|
||||
pool[i] = byte(rng.Uint32())
|
||||
}
|
||||
// Random pool large enough for the longest case + alignment slop.
|
||||
pool := make([]byte, 4096+padFront)
|
||||
for i := range pool {
|
||||
pool[i] = byte(rng.Uint32())
|
||||
}
|
||||
|
||||
seeds := []uint16{0, 0x0001, 0xabcd, 0xffff, 0x1234, 0xfedc}
|
||||
offsets := []int{0, 1, 2, 3, 4, 5, 7, 8, 15, 16}
|
||||
seeds := []uint16{0, 0x0001, 0xabcd, 0xffff, 0x1234, 0xfedc}
|
||||
offsets := []int{0, 1, 2, 3, 4, 5, 7, 8, 15, 16}
|
||||
|
||||
for length := 0; length <= 4096; length++ {
|
||||
for _, seed := range seeds {
|
||||
for _, off := range offsets {
|
||||
if off+length > len(pool) {
|
||||
continue
|
||||
}
|
||||
buf := pool[off : off+length]
|
||||
want := gvisorchecksum.Checksum(buf, seed)
|
||||
got := impl.fn(buf, seed)
|
||||
if got != want {
|
||||
t.Fatalf("len=%d off=%d seed=%#x: got %#04x want %#04x",
|
||||
length, off, seed, got, want)
|
||||
}
|
||||
}
|
||||
for length := 0; length <= 4096; length++ {
|
||||
for _, seed := range seeds {
|
||||
for _, off := range offsets {
|
||||
if off+length > len(pool) {
|
||||
continue
|
||||
}
|
||||
buf := pool[off : off+length]
|
||||
want := gvisorchecksum.Checksum(buf, seed)
|
||||
got := Checksum(buf, seed)
|
||||
if got != want {
|
||||
t.Fatalf("len=%d off=%d seed=%#x: got %#04x want %#04x",
|
||||
length, off, seed, got, want)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -78,28 +46,23 @@ func TestChecksumMatchesGvisor(t *testing.T) {
|
||||
// historically tripped up checksum implementations: all-zero, all-0xff,
|
||||
// alternating, and ascending sequences.
|
||||
func TestChecksumPatternedBuffers(t *testing.T) {
|
||||
for _, impl := range implsUnderTest() {
|
||||
t.Run(impl.name, func(t *testing.T) {
|
||||
requireAvailable(t, impl)
|
||||
for length := 0; length <= 256; length++ {
|
||||
patterns := map[string][]byte{
|
||||
"zeros": make([]byte, length),
|
||||
"ones": bytes(length, 0xff),
|
||||
"alternating": pattern(length, []byte{0xa5, 0x5a}),
|
||||
"ascending": ascending(length),
|
||||
}
|
||||
for name, buf := range patterns {
|
||||
for _, seed := range []uint16{0, 0xffff, 0x8000} {
|
||||
want := gvisorchecksum.Checksum(buf, seed)
|
||||
got := impl.fn(buf, seed)
|
||||
if got != want {
|
||||
t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x",
|
||||
name, length, seed, got, want)
|
||||
}
|
||||
}
|
||||
for length := 0; length <= 256; length++ {
|
||||
patterns := map[string][]byte{
|
||||
"zeros": make([]byte, length),
|
||||
"ones": bytes(length, 0xff),
|
||||
"alternating": pattern(length, []byte{0xa5, 0x5a}),
|
||||
"ascending": ascending(length),
|
||||
}
|
||||
for name, buf := range patterns {
|
||||
for _, seed := range []uint16{0, 0xffff, 0x8000} {
|
||||
want := gvisorchecksum.Checksum(buf, seed)
|
||||
got := Checksum(buf, seed)
|
||||
if got != want {
|
||||
t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x",
|
||||
name, length, seed, got, want)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -135,41 +98,36 @@ func ascending(n int) []byte {
|
||||
// and k=1 (one main loop iter, then tail). It's explicit coverage for
|
||||
// payload sizes that are odd, not divisible by 4, by 8, or by 32.
|
||||
func TestChecksumTailPaths(t *testing.T) {
|
||||
for _, impl := range implsUnderTest() {
|
||||
t.Run(impl.name, func(t *testing.T) {
|
||||
requireAvailable(t, impl)
|
||||
rng := rand.New(rand.NewPCG(42, 17))
|
||||
const padFront = 16
|
||||
const maxK = 8
|
||||
rng := rand.New(rand.NewPCG(42, 17))
|
||||
const padFront = 16
|
||||
const maxK = 8
|
||||
|
||||
pool := make([]byte, 64*maxK+padFront+64)
|
||||
for i := range pool {
|
||||
pool[i] = byte(rng.Uint32())
|
||||
}
|
||||
pool := make([]byte, 64*maxK+padFront+64)
|
||||
for i := range pool {
|
||||
pool[i] = byte(rng.Uint32())
|
||||
}
|
||||
|
||||
seeds := []uint16{0, 0xffff, 0xabcd}
|
||||
offsets := []int{0, 1, 3, 7, 15} // mix of aligned and odd starts
|
||||
seeds := []uint16{0, 0xffff, 0xabcd}
|
||||
offsets := []int{0, 1, 3, 7, 15} // mix of aligned and odd starts
|
||||
|
||||
for k := 0; k <= maxK; k++ {
|
||||
for tail := 0; tail < 64; tail++ {
|
||||
length := 64*k + tail
|
||||
for _, seed := range seeds {
|
||||
for _, off := range offsets {
|
||||
if off+length > len(pool) {
|
||||
continue
|
||||
}
|
||||
buf := pool[off : off+length]
|
||||
want := gvisorchecksum.Checksum(buf, seed)
|
||||
got := impl.fn(buf, seed)
|
||||
if got != want {
|
||||
t.Fatalf("k=%d tail=%d (len=%d) off=%d seed=%#x: got %#04x want %#04x",
|
||||
k, tail, length, off, seed, got, want)
|
||||
}
|
||||
}
|
||||
for k := 0; k <= maxK; k++ {
|
||||
for tail := 0; tail < 64; tail++ {
|
||||
length := 64*k + tail
|
||||
for _, seed := range seeds {
|
||||
for _, off := range offsets {
|
||||
if off+length > len(pool) {
|
||||
continue
|
||||
}
|
||||
buf := pool[off : off+length]
|
||||
want := gvisorchecksum.Checksum(buf, seed)
|
||||
got := Checksum(buf, seed)
|
||||
if got != want {
|
||||
t.Fatalf("k=%d tail=%d (len=%d) off=%d seed=%#x: got %#04x want %#04x",
|
||||
k, tail, length, off, seed, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,11 +0,0 @@
|
||||
package checksum
|
||||
|
||||
// archImpls exposes every hand-written implementation on this architecture
|
||||
// so the tests exercise them directly, independent of what the public
|
||||
// Checksum dispatches to on the running CPU. Without this, running the
|
||||
// suite on a non-AVX2 machine compared gvisor against itself and left the
|
||||
// assembly untested — silently. available=false makes the test skip loudly
|
||||
// instead.
|
||||
var archImpls = []archImpl{
|
||||
{name: "avx2", fn: checksumAVX2, available: hasAVX2},
|
||||
}
|
||||
@@ -1,8 +0,0 @@
|
||||
package checksum
|
||||
|
||||
// archImpls exposes every hand-written implementation on this architecture
|
||||
// for direct testing; see export_amd64_test.go for the rationale. NEON is
|
||||
// mandatory in armv8, so it is always available.
|
||||
var archImpls = []archImpl{
|
||||
{name: "neon", fn: checksumNEON, available: true},
|
||||
}
|
||||
@@ -1,7 +0,0 @@
|
||||
//go:build !amd64 && !arm64
|
||||
|
||||
package checksum
|
||||
|
||||
// No hand-written implementations on this architecture; the dispatcher is
|
||||
// pure gvisor and there is nothing separate to test.
|
||||
var archImpls []archImpl
|
||||
@@ -9,9 +9,10 @@ import (
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// blockOn parks the calling goroutine until fd is ready or shutdownFd signals teardown.
|
||||
// (events is POLLIN for reads, POLLOUT for writes)
|
||||
// It builds the pollfd array on the stack every call, so concurrent callers on the same Queue never share Revents storage.
|
||||
// blockOn parks the calling goroutine until fd is ready (events is POLLIN for
|
||||
// reads, POLLOUT for writes) or shutdownFd signals teardown. It builds the
|
||||
// pollfd array on the stack every call, so concurrent callers on the same
|
||||
// Queue never share Revents storage.
|
||||
//
|
||||
// Returns os.ErrClosed when shutdown was signaled (POLLIN on shutdownFd)
|
||||
// or either fd reported a problem condition (POLLHUP|POLLNVAL|POLLERR).
|
||||
|
||||
@@ -7,7 +7,6 @@ import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"sync/atomic"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
@@ -18,17 +17,18 @@ type offloadQueueSet struct {
|
||||
// pqi is exactly the same as pq, but stored as the interface type
|
||||
pqi []Queue
|
||||
shutdownFd int
|
||||
// usoEnabled is true when newTun successfully negotiated TUN_F_USO4|6 with the kernel.
|
||||
// Queues created by Add inherit this and surface it via Offload.USOSupported so coalescers can gate USO emission.
|
||||
// usoEnabled is true when newTun successfully negotiated TUN_F_USO4|6
|
||||
// with the kernel. Queues created by Add inherit this and surface it
|
||||
// via Offload.USOSupported so coalescers can gate USO emission.
|
||||
usoEnabled bool
|
||||
closed atomic.Bool
|
||||
// l is handed to each queue for its bad-vnet-header drop logging.
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
// NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do TSO segmentation.
|
||||
// usoEnabled tells downstream queues whether the kernel agreed to deliver/accept GSO_UDP_L4 superpackets.
|
||||
func NewOffloadQueueSet(usoEnabled bool, l *slog.Logger) (QueueSet, error) {
|
||||
// NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do
|
||||
// TSO segmentation in userspace. usoEnabled tells downstream queues whether
|
||||
// the kernel agreed to deliver/accept GSO_UDP_L4 superpackets — coalescers
|
||||
// should fall back to per-packet writes when this is false.
|
||||
func NewOffloadQueueSet(usoEnabled bool) (QueueSet, error) {
|
||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||
@@ -39,7 +39,6 @@ func NewOffloadQueueSet(usoEnabled bool, l *slog.Logger) (QueueSet, error) {
|
||||
pqi: []Queue{},
|
||||
shutdownFd: shutdownFd,
|
||||
usoEnabled: usoEnabled,
|
||||
l: l,
|
||||
}
|
||||
|
||||
return out, nil
|
||||
@@ -50,10 +49,7 @@ func (c *offloadQueueSet) Queues() []Queue {
|
||||
}
|
||||
|
||||
func (c *offloadQueueSet) Add(fd int) error {
|
||||
if c.closed.Load() {
|
||||
return errors.New("queue set already closed")
|
||||
}
|
||||
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled, c.l)
|
||||
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -77,21 +73,23 @@ func (c *offloadQueueSet) Close() error {
|
||||
|
||||
errs := []error{}
|
||||
|
||||
// Signal all readers blocked in poll to wake up and exit.
|
||||
// They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
|
||||
// Signal all readers blocked in poll to wake up and exit. They observe
|
||||
// POLLIN on the shutdown eventfd and return os.ErrClosed.
|
||||
if err := c.wakeForShutdown(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
|
||||
// Close the per-queue tun fds; this also unblocks any in-flight reads.
|
||||
// The per-queue Close deliberately leaves shutdownFd alone - it belongs
|
||||
// to this container.
|
||||
for _, x := range c.pq {
|
||||
if err := x.Close(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Close the shutdown eventfd last: every reader's pollfd set references it,
|
||||
// so it must outlive the wake + per-queue teardown above.
|
||||
// Close the shutdown eventfd last: every reader's pollfd set references
|
||||
// it, so it must outlive the wake + per-queue teardown above.
|
||||
if err := unix.Close(c.shutdownFd); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
|
||||
@@ -40,9 +40,6 @@ func (c *pollQueueSet) Queues() []Queue {
|
||||
}
|
||||
|
||||
func (c *pollQueueSet) Add(fd int) error {
|
||||
if c.closed.Load() {
|
||||
return errors.New("queue set already closed")
|
||||
}
|
||||
x, err := newPoll(fd, c.shutdownFd)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -67,21 +64,23 @@ func (c *pollQueueSet) Close() error {
|
||||
|
||||
errs := []error{}
|
||||
|
||||
// Signal all readers blocked in poll to wake up and exit.
|
||||
// They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
|
||||
// Wake any reader blocked in poll so it observes POLLIN on the shutdown
|
||||
// eventfd and returns os.ErrClosed.
|
||||
if err := c.wakeForShutdown(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
|
||||
// Close the per-queue tun fds; this also unblocks any in-flight reads.
|
||||
// The per-queue Close deliberately leaves shutdownFd alone - it belongs
|
||||
// to this container.
|
||||
for _, x := range c.pq {
|
||||
if err := x.Close(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Close the shutdown eventfd last: every reader's pollfd set references it,
|
||||
// so it must outlive the wake + per-queue teardown above.
|
||||
// Close the shutdown eventfd last: every reader's pollfd set references
|
||||
// it, so it must outlive the wake + per-queue teardown above.
|
||||
if err := unix.Close(c.shutdownFd); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
//go:build !linux || android
|
||||
//go:build !linux || android || e2e_testing
|
||||
|
||||
package tio
|
||||
|
||||
@@ -8,6 +8,12 @@ func protoFromGSOType(_ uint8) (GSOProto, error) {
|
||||
return 0, fmt.Errorf("GSO unsupported")
|
||||
}
|
||||
|
||||
// SegmentSuperpacket invokes fn once per segment of pkt. On non-Linux
|
||||
// builds (and Android/e2e_testing) this package does not provide a Queue
|
||||
// implementation, so any caller that does construct a Packet here can only
|
||||
// be operating on non-superpacket bytes and the stub forwards them
|
||||
// directly. A non-zero GSO field is a programming error from the caller
|
||||
// and returns an explicit error rather than silently misbehaving.
|
||||
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
||||
if pkt.GSO.IsSuperpacket() {
|
||||
return fmt.Errorf("tio: GSO superpacket on platform without segmentation support")
|
||||
|
||||
@@ -4,8 +4,9 @@ import "io"
|
||||
|
||||
// singleQueue adapts a legacy one-datagram-per-Read source into a Queue.
|
||||
// Read fills a private scratch buffer and returns exactly one Packet whose
|
||||
// Bytes borrow from that buffer, valid only until the next Read, per the Queue contract.
|
||||
// Single-reader like every Queue; Write is exactly as safe for concurrent use as the underlying source's Write.
|
||||
// Bytes borrow from that buffer, valid only until the next Read, per the
|
||||
// Queue contract. Single-reader like every Queue; Write is exactly as safe
|
||||
// for concurrent use as the underlying source's Write.
|
||||
type singleQueue struct {
|
||||
rw io.ReadWriter
|
||||
closer io.Closer // nil: Close is a no-op (the source is shared and owned elsewhere)
|
||||
@@ -13,9 +14,9 @@ type singleQueue struct {
|
||||
ret [1]Packet
|
||||
}
|
||||
|
||||
// NewSingleQueue wraps a one-datagram-per-Read ReadWriteCloser (a legacy tun device) into a Queue.
|
||||
// bufSize is the per-queue read scratch size and must be at least the largest datagram the source can return.
|
||||
// Close closes rwc.
|
||||
// NewSingleQueue wraps a one-datagram-per-Read ReadWriteCloser (a legacy tun
|
||||
// device) into a Queue. bufSize is the per-queue read scratch size and must
|
||||
// be at least the largest datagram the source can return. Close closes rwc.
|
||||
func NewSingleQueue(rwc io.ReadWriteCloser, bufSize int) Queue {
|
||||
return &singleQueue{rw: rwc, closer: rwc, buf: make([]byte, bufSize)}
|
||||
}
|
||||
|
||||
+85
-60
@@ -13,72 +13,79 @@ type QueueSet interface {
|
||||
Add(fd int) error
|
||||
}
|
||||
|
||||
// Capabilities advertises which kernel offload features a Queue successfully negotiated.
|
||||
// Callers consult this to decide which coalescers to wire onto the write path.
|
||||
// Capabilities advertises which kernel offload features a Queue
|
||||
// successfully negotiated. Callers consult this to decide which coalescers
|
||||
// to wire onto the write path — a Queue without TSO can't usefully accept a
|
||||
// TCPCoalescer, and a Queue without USO can't accept a UDPCoalescer.
|
||||
type Capabilities struct {
|
||||
// TSO means the FD was opened with IFF_VNET_HDR and the kernel agreed to TUN_F_TSO4|TSO6,
|
||||
// and WriteGSO with GSOProtoTCP is safe.
|
||||
// TSO means the FD was opened with IFF_VNET_HDR and the kernel agreed
|
||||
// to TUN_F_TSO4|TSO6 — i.e. WriteGSO with GSOProtoTCP is safe.
|
||||
TSO bool
|
||||
// USO means the kernel additionally agreed to TUN_F_USO4|USO6,
|
||||
// so WriteGSO with GSOProtoUDP is safe. Linux ≥ 6.2.
|
||||
// USO means the kernel additionally agreed to TUN_F_USO4|USO6, so
|
||||
// WriteGSO with GSOProtoUDP is safe. Linux ≥ 6.2.
|
||||
USO bool
|
||||
}
|
||||
|
||||
// Queue is a readable/writable Poll queue.
|
||||
// Concurrency contract: a single read goroutine drives Read; plain Write is safe for concurrent callers;
|
||||
// Queue is a readable/writable Poll queue. Concurrency contract: a single
|
||||
// read goroutine drives Read; plain Write is safe for concurrent callers;
|
||||
// WriteGSO (on Queues that implement GSOWriter) is single-writer per queue.
|
||||
//
|
||||
// Close on an individual Queue does NOT unblock a Read parked in poll — closing an fd
|
||||
// never wakes its pollers. Orderly teardown goes through the owning QueueSet's Close,
|
||||
// which first signals a shared shutdown eventfd every reader polls alongside its own fd.
|
||||
// That eventfd is a set-wide kill switch: once signaled, every Queue in the set returns
|
||||
// os.ErrClosed from Read, so it cannot be used to stop a single Queue.
|
||||
type Queue interface {
|
||||
io.Closer
|
||||
|
||||
// Read returns one or more packets.
|
||||
// The returned Packet.Bytes slices are borrowed from the Queue's internal buffer and are only valid
|
||||
// until the next Read or Close on this Queue.
|
||||
// A Packet may carry a GSO/USO superpacket (see GSOInfo)
|
||||
// Single-reader only: not safe for concurrent Reads (it reuses per-queue rx scratch each call).
|
||||
// Read returns one or more packets. The returned Packet.Bytes slices
|
||||
// are borrowed from the Queue's internal buffer and are only valid
|
||||
// until the next Read or Close on this Queue - callers must encrypt
|
||||
// or copy each slice before the next call. A Packet may carry a
|
||||
// GSO/USO superpacket (see GSOInfo); when GSO.IsSuperpacket() is
|
||||
// true the caller must segment Bytes before treating it as a single
|
||||
// IP datagram. Single-reader only: not safe for concurrent Reads (it
|
||||
// reuses per-queue rx scratch each call).
|
||||
Read() ([]Packet, error)
|
||||
|
||||
// Write emits a single packet on the plaintext (outside→inside) delivery path.
|
||||
// Safe for concurrent use.
|
||||
// Write emits a single packet on the plaintext (outside→inside)
|
||||
// delivery path. Safe for concurrent use.
|
||||
Write(p []byte) (int, error)
|
||||
}
|
||||
|
||||
// Packet is the unit Queue.Read returns.
|
||||
// Bytes points into the queue's internal buffer and is only valid until the next Read or Close on the queue that produced it.
|
||||
// GSO is the zero value for an already-segmented IP datagram;
|
||||
// when non-zero it describes a kernel-supplied TSO/USO superpacket the caller must segment before consuming.
|
||||
// Packet is the unit Queue.Read returns. Bytes points into the queue's
|
||||
// internal buffer and is only valid until the next Read or Close on the
|
||||
// queue that produced it. GSO is the zero value for an already-segmented
|
||||
// IP datagram; when non-zero it describes a kernel-supplied TSO/USO
|
||||
// superpacket the caller must segment before consuming.
|
||||
type Packet struct {
|
||||
Bytes []byte
|
||||
GSO GSOInfo
|
||||
}
|
||||
|
||||
// GSOInfo describes a kernel-supplied superpacket sitting in Packet.Bytes.
|
||||
// The zero value means Bytes is one regular IP datagram and no segmentation is required.
|
||||
// The zero value means "not a superpacket" — Bytes is one regular IP
|
||||
// datagram and no segmentation is required.
|
||||
type GSOInfo struct {
|
||||
// Size is the GSO segment size: max payload bytes per segment
|
||||
// (== TCP MSS for TSO, == UDP payload chunk for USO). Zero means not a superpacket.
|
||||
// (== TCP MSS for TSO, == UDP payload chunk for USO). Zero means
|
||||
// not a superpacket.
|
||||
Size uint16
|
||||
// HdrLen is the total L3+L4 header length within Bytes (already corrected via correctHdrLen, so safe to slice on).
|
||||
// HdrLen is the total L3+L4 header length within Bytes (already
|
||||
// corrected via correctHdrLen, so safe to slice on).
|
||||
HdrLen uint16
|
||||
// CsumStart is the L4 header offset inside Bytes (== L3 header length).
|
||||
// CsumStart is the L4 header offset inside Bytes (== L3 header
|
||||
// length).
|
||||
CsumStart uint16
|
||||
// Proto picks the L4 protocol (TCP or UDP) so the segmenter knows which checksum/header layout to apply.
|
||||
// Proto picks the L4 protocol (TCP or UDP) so the segmenter knows
|
||||
// which checksum/header layout to apply.
|
||||
Proto GSOProto
|
||||
}
|
||||
|
||||
// IsSuperpacket reports whether g describes a multi-segment GSO/USO
|
||||
// superpacket that needs segmentation before its bytes can be encrypted and sent on the wire.
|
||||
// superpacket that needs segmentation before its bytes can be encrypted
|
||||
// and sent on the wire.
|
||||
func (g GSOInfo) IsSuperpacket() bool { return g.Size > 0 }
|
||||
|
||||
// Clone returns a Packet whose Bytes is a freshly allocated copy of p.Bytes,
|
||||
// safe to retain past the next Read or Close on the originating Queue.
|
||||
// GSO metadata is copied verbatim.
|
||||
// Use this only when a caller needs the data to outlive the borrowed-slice contract.
|
||||
// GSO metadata is copied verbatim. Use this only when a caller genuinely
|
||||
// needs to outlive the borrowed-slice contract — the hot path reads should
|
||||
// continue to consume the borrow synchronously to avoid the allocation.
|
||||
func (p Packet) Clone() Packet {
|
||||
if p.Bytes == nil {
|
||||
return p
|
||||
@@ -88,60 +95,78 @@ func (p Packet) Clone() Packet {
|
||||
return Packet{Bytes: cp, GSO: p.GSO}
|
||||
}
|
||||
|
||||
// CapsProvider is an optional interface implemented by Queues that negotiate kernel offload features at open time.
|
||||
// Callers pick a write-path coalescer based on the result.
|
||||
// Queues that don't implement it are treated as having no offload capability.
|
||||
// CapsProvider is an optional interface implemented by Queues that
|
||||
// successfully negotiated kernel offload features at open time. Callers
|
||||
// pick a write-path coalescer based on the result. Queues that don't
|
||||
// implement it are treated as having no offload capability — callers must
|
||||
// fall back to plain per-packet writes.
|
||||
type CapsProvider interface {
|
||||
Capabilities() Capabilities
|
||||
}
|
||||
|
||||
// GSOProto selects the L4 protocol for a GSO superpacket.
|
||||
// Determines which VIRTIO_NET_HDR_GSO_* type the writer stamps and which checksum offset
|
||||
// QueueCapabilities returns q's negotiated offload capabilities, or the
|
||||
// zero value when q does not advertise any.
|
||||
func QueueCapabilities(q Queue) Capabilities {
|
||||
if cp, ok := q.(CapsProvider); ok {
|
||||
return cp.Capabilities()
|
||||
}
|
||||
return Capabilities{}
|
||||
}
|
||||
|
||||
// GSOProto selects the L4 protocol for a GSO superpacket. Determines which
|
||||
// VIRTIO_NET_HDR_GSO_* type the writer stamps and which checksum offset
|
||||
// inside the transport header virtio NEEDS_CSUM expects.
|
||||
type GSOProto uint8
|
||||
|
||||
const (
|
||||
GSOProtoUnknown GSOProto = iota
|
||||
GSOProtoTCP
|
||||
GSOProtoTCP GSOProto = iota
|
||||
GSOProtoUDP
|
||||
)
|
||||
|
||||
// GSOWriter is implemented by Queues that can emit a TCP or UDP superpacket
|
||||
// assembled from a header prefix plus one or more borrowed payload fragments,
|
||||
// in a single vectored write (writev with a leading virtio_net_hdr).
|
||||
// This lets the coalescer avoid copying payload bytes between the caller's decrypt buffer and the TUN.
|
||||
// Backends without GSO support do not implement this interface and coalescing is skipped.
|
||||
// assembled from a header prefix plus one or more borrowed payload
|
||||
// fragments, in a single vectored write (writev with a leading
|
||||
// virtio_net_hdr). This lets the coalescer avoid copying payload bytes
|
||||
// between the caller's decrypt buffer and the TUN. Backends without GSO
|
||||
// support do not implement this interface and coalescing is skipped.
|
||||
//
|
||||
// hdr contains the IPv4/IPv6 header prefix (mutable: callers will have filled in total length and IP csum).
|
||||
// transportHdr is the TCP or UDP header
|
||||
// (mutable: the L4 checksum field must hold the pseudo-header partial, single-fold not inverted, per virtio NEEDS_CSUM semantics).
|
||||
// pays are non-overlapping payload fragments whose concatenation is the full superpacket payload.
|
||||
// They are read-only from the writer's perspective and must remain valid until the call returns.
|
||||
// Every segment in pays except possibly the last must be exactly the same size.
|
||||
// proto picks the L4 protocol so the writer knows which gsoType / CsumOffset to set.
|
||||
// hdr contains the IPv4/IPv6 header prefix (mutable - callers will have
|
||||
// filled in total length and IP csum). transportHdr is the TCP or UDP
|
||||
// header (mutable - the L4 checksum field must hold the pseudo-header
|
||||
// partial, single-fold not inverted, per virtio NEEDS_CSUM semantics).
|
||||
// pays are non-overlapping payload fragments whose concatenation is the
|
||||
// full superpacket payload; they are read-only from the writer's
|
||||
// perspective and must remain valid until the call returns. Every segment
|
||||
// in pays except possibly the last is exactly the same size. proto picks
|
||||
// the L4 protocol so the writer knows which GSOType / CsumOffset to set.
|
||||
//
|
||||
// Callers should also consult CapsProvider (via SupportsGSO) for the per-protocol negotiated capability:
|
||||
// USO may not have been negotiated even when TSO was.
|
||||
// Callers should also consult CapsProvider (via SupportsGSO or
|
||||
// QueueCapabilities) for the per-protocol negotiated capability; an
|
||||
// implementation of GSOWriter is necessary but not sufficient since USO
|
||||
// may not have been negotiated even when TSO was.
|
||||
type GSOWriter interface {
|
||||
io.Writer
|
||||
CapsProvider
|
||||
WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error
|
||||
}
|
||||
|
||||
// SupportsGSO reports whether w implements GSOWriter and the underlying
|
||||
// queue advertises the negotiated capability for `want`.
|
||||
func SupportsGSO(w io.Writer, want GSOProto) (GSOWriter, bool) {
|
||||
// queue advertises the negotiated capability for `want`. A writer that
|
||||
// implements GSOWriter but not CapsProvider is treated as permissive
|
||||
// (used by tests and fakes that don't negotiate).
|
||||
func SupportsGSO(w any, want GSOProto) (GSOWriter, bool) {
|
||||
gw, ok := w.(GSOWriter)
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
caps := gw.Capabilities()
|
||||
cp, ok := w.(CapsProvider)
|
||||
if !ok {
|
||||
return gw, true
|
||||
}
|
||||
caps := cp.Capabilities()
|
||||
switch want {
|
||||
case GSOProtoTCP:
|
||||
return gw, caps.TSO
|
||||
case GSOProtoUDP:
|
||||
return gw, caps.USO
|
||||
default:
|
||||
return gw, false
|
||||
}
|
||||
return gw, false
|
||||
}
|
||||
|
||||
+150
-153
@@ -4,7 +4,6 @@
|
||||
package tio
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
@@ -18,56 +17,71 @@ import (
|
||||
"github.com/slackhq/nebula/overlay/tio/virtio"
|
||||
)
|
||||
|
||||
const maxSuperpacketLen = 65535
|
||||
|
||||
// tunRxBufSize is the per-Read worst-case footprint inside rxBuf: one kernel-supplied packet body, which is at most ~64 KiB.
|
||||
// tunRxBufSize is the per-Read worst-case footprint inside rxBuf: one
|
||||
// kernel-supplied packet body, which is at most ~64 KiB (tunReadBufSize).
|
||||
// Segmentation happens at encrypt time on a per-routine MTU-sized scratch
|
||||
// (see SegmentSuperpacket), so rxBuf only holds raw kernel-supplied bytes.
|
||||
// We round up to give margin for the drain headroom check below.
|
||||
// We round up to give comfortable margin for the drain headroom check
|
||||
// below.
|
||||
const tunRxBufSize = 64 * 1024
|
||||
|
||||
// tunRxBufCap is the total size we allocate for the per-reader rx buffer.
|
||||
// Each drain iteration consumes up to tunRxBufSize of headroom for the kernel-supplied bytes.
|
||||
// Sized to eight such iterations so a single poll wake can drain several TSO/USO superpackets under bulk load,
|
||||
// amortizing the wake and giving the sendmmsg planner longer same-destination runs.
|
||||
// Hold latency stays bounded because listenIn flushes its send batch incrementally rather than only at end-of-drain.
|
||||
// tunRxBufCap is the total size we allocate for the per-reader rx
|
||||
// buffer. With reads landing directly in rxBuf, each drain iteration
|
||||
// consumes up to tunRxBufSize of headroom for the kernel-supplied bytes.
|
||||
// Sized to eight such iterations so a single poll wake can drain several
|
||||
// TSO/USO superpackets under bulk load, amortizing the wake and giving
|
||||
// the sendmmsg planner longer same-destination runs. Hold latency stays
|
||||
// bounded because listenIn flushes its send batch incrementally rather
|
||||
// than only at end-of-drain.
|
||||
const tunRxBufCap = tunRxBufSize * 8
|
||||
|
||||
// tunDrainCap caps how many packets a single Read will accumulate via the post-wake drain loop.
|
||||
// Sized to soak up a burst of small ACKs while bounding how much work a single caller holds before handing off.
|
||||
// tunDrainCap caps how many packets a single Read will accumulate via
|
||||
// the post-wake drain loop. Sized to soak up a burst of small ACKs while
|
||||
// bounding how much work a single caller holds before handing off.
|
||||
const tunDrainCap = 64
|
||||
|
||||
// gsoMaxIovs caps the iovec budget WriteGSO assembles per call:
|
||||
// 3 fixed entries (virtio_net_hdr, IP hdr, transport hdr), plus up to gsoMaxIovs-3 payload fragments.
|
||||
// Sized comfortably above the typical kernel GSO segment cap (Linux UDP_GRO is 64)
|
||||
// so realistic coalesced bursts never touch the limit.
|
||||
// iovecs are tiny (16 bytes), so the entire scratch is 4 KiB.
|
||||
// WriteGSO returns an error rather than reallocating when a caller exceeds this budget.
|
||||
// gsoMaxIovs caps the iovec budget WriteGSO assembles per call: 3 fixed
|
||||
// entries (virtio_net_hdr, IP hdr, transport hdr) plus up to gsoMaxIovs-3
|
||||
// payload fragments. Sized comfortably above the typical kernel GSO
|
||||
// segment cap (Linux UDP_GRO is 64) so realistic coalesced bursts never
|
||||
// touch the limit. iovecs are tiny (16 bytes), so the entire scratch is
|
||||
// 4 KiB — fine to keep resident on every queue. WriteGSO returns an error
|
||||
// rather than reallocating when a caller exceeds this budget.
|
||||
const gsoMaxIovs = 256
|
||||
|
||||
// validVnetHdr is the 10-byte virtio_net_hdr we prepend to every non-GSO TUN write.
|
||||
// Only flag set is VIRTIO_NET_HDR_F_DATA_VALID. Note the tun write path
|
||||
// (__virtio_net_hdr_to_skb) ignores this bit — only the virtio-net driver's RX
|
||||
// helper honors it — so packets land CHECKSUM_NONE and the stack verifies the
|
||||
// L4 checksum anyway. What matters here is what the header does NOT say:
|
||||
// no NEEDS_CSUM, so the kernel is never asked to finish a checksum.
|
||||
// All packets that reach the plain Write paths already carry a valid L4 checksum.
|
||||
// validVnetHdr is the 10-byte virtio_net_hdr we prepend to every non-GSO TUN
|
||||
// write. Only flag set is VIRTIO_NET_HDR_F_DATA_VALID, which marks the skb
|
||||
// CHECKSUM_UNNECESSARY so the receiving network stack skips L4 checksum
|
||||
// verification. All packets that reach the plain Write paths already carry
|
||||
// a valid L4 checksum (either supplied by a remote peer whose ciphertext we
|
||||
// AEAD-authenticated, produced by segmentTCPYield/segmentUDPYield during
|
||||
// superpacket segmentation, or built locally by CreateRejectPacket), so
|
||||
// trusting them is safe.
|
||||
var validVnetHdr = [virtio.Size]byte{unix.VIRTIO_NET_HDR_F_DATA_VALID}
|
||||
|
||||
// Offload wraps a TUN file descriptor with poll-based reads. The FD provided will be changed to non-blocking.
|
||||
// A shared eventfd allows Close to wake all readers blocked in poll.
|
||||
//
|
||||
// Field order is deliberate: the read-mostly fds and the writer-owned GSO scratch fill
|
||||
// the first cache line, and the state the reader mutates per packet (rxOff, pending,
|
||||
// readIovs) all sits after it, so per-packet reader stores never invalidate the line
|
||||
// concurrent Write callers load fd from.
|
||||
type Offload struct {
|
||||
fd int
|
||||
shutdownFd int
|
||||
closed atomic.Bool
|
||||
rxBuf []byte // backing store for kernel-handed packets read this drain
|
||||
rxOff int // cursor into rxBuf for the current Read drain
|
||||
pending []Packet // packets returned from the most recent Read
|
||||
|
||||
// readVnetScratch holds the 10-byte virtio_net_hdr split off the front of
|
||||
// every TUN read via readv(2). Decoupling the header from the packet body
|
||||
// lets us read the body directly into rxBuf at the current rxOff with
|
||||
// no userspace copy on the GSO_NONE fast path.
|
||||
readVnetScratch [virtio.Size]byte
|
||||
// readIovs is the readv(2) iovec scratch wired once at construction —
|
||||
// iovec[0] points at readVnetScratch; iovec[1].Base/Len is updated per
|
||||
// read to address the current rxBuf slot.
|
||||
readIovs [2]unix.Iovec
|
||||
|
||||
// usoEnabled records whether the kernel agreed to TUN_F_USO* on this FD,
|
||||
// so writers can decide whether emitting GSO_UDP_L4 superpackets is safe.
|
||||
usoEnabled bool
|
||||
closed atomic.Bool
|
||||
|
||||
// gsoHdrBuf is a per-queue 10-byte scratch for the virtio_net_hdr emitted
|
||||
// by WriteGSO. Kept separate from the read-only package-level validVnetHdr
|
||||
@@ -78,27 +92,9 @@ type Offload struct {
|
||||
// gsoMaxIovs at construction; never grown. WriteGSO returns an error
|
||||
// (and drops the call) if a caller hands it more fragments than fit.
|
||||
gsoIovs []unix.Iovec
|
||||
|
||||
rxBuf []byte // backing store for kernel-handed packets read this drain
|
||||
rxOff int // cursor into rxBuf for the current Read drain
|
||||
pending []Packet // packets returned from the most recent Read
|
||||
|
||||
// readVnetScratch holds the 10-byte virtio_net_hdr split off the front of
|
||||
// every TUN read via readv(2). Decoupling the header from the packet body
|
||||
// lets us read the body directly into rxBuf at the current rxOff with
|
||||
// no userspace copy on the GSO_NONE fast path.
|
||||
readVnetScratch [virtio.Size]byte
|
||||
// readIovs is the readv(2) iovec scratch wired once at construction,
|
||||
// iovec[0] points at readVnetScratch
|
||||
// iovec[1].Base/Len is updated per read to address the current rxBuf slot.
|
||||
readIovs [2]unix.Iovec
|
||||
|
||||
// l is only consulted on the rare bad-vnet-header drop path; it lives
|
||||
// after the hot state on purpose. May be nil (tests); drops go unlogged then.
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
func newOffload(fd int, shutdownFd int, usoEnabled bool, l *slog.Logger) (*Offload, error) {
|
||||
func newOffload(fd int, shutdownFd int, usoEnabled bool) (*Offload, error) {
|
||||
if err := unix.SetNonblock(fd, true); err != nil {
|
||||
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
||||
}
|
||||
@@ -108,7 +104,6 @@ func newOffload(fd int, shutdownFd int, usoEnabled bool, l *slog.Logger) (*Offlo
|
||||
shutdownFd: shutdownFd,
|
||||
usoEnabled: usoEnabled,
|
||||
closed: atomic.Bool{},
|
||||
l: l,
|
||||
|
||||
rxBuf: make([]byte, tunRxBufCap),
|
||||
gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs),
|
||||
@@ -133,18 +128,26 @@ func (r *Offload) blockOnWrite() error {
|
||||
return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLOUT)
|
||||
}
|
||||
|
||||
// readPacket issues a single readv(2), splitting the virtio_net_hdr off into readVnetScratch
|
||||
// and reading the packet body directly into rxBuf at the current rxOff.
|
||||
// Returns the body length (zero virtio header bytes, just the IP packet/superpacket).
|
||||
// block controls whether EAGAIN is retried via poll: the initial read of a drain blocks; subsequent drain reads do not.
|
||||
// readPacket issues a single readv(2) splitting the virtio_net_hdr off
|
||||
// into readVnetScratch and reading the packet body directly into rxBuf at
|
||||
// the current rxOff. Returns the body length (zero virtio header bytes,
|
||||
// just the IP packet/superpacket). block controls whether EAGAIN is
|
||||
// retried via poll: the initial read of a drain blocks; subsequent drain
|
||||
// reads do not.
|
||||
//
|
||||
// The body iovec capacity is always tunReadBufSize; callers (the Read
|
||||
// drain loop) gate entry on tunRxBufCap-rxOff >= tunRxBufSize, sized to
|
||||
// hold one worst-case kernel-supplied packet body. Without that gate the
|
||||
// body iovec could be smaller than the next inbound packet and the
|
||||
// kernel would truncate.
|
||||
func (r *Offload) readPacket(block bool) (int, error) {
|
||||
for {
|
||||
r.readIovs[1].Base = &r.rxBuf[r.rxOff]
|
||||
r.readIovs[1].SetLen(len(r.rxBuf) - r.rxOff)
|
||||
r.readIovs[1].SetLen(tunReadBufSize)
|
||||
n, _, errno := syscall.Syscall(unix.SYS_READV, uintptr(r.fd), uintptr(unsafe.Pointer(&r.readIovs[0])), uintptr(len(r.readIovs)))
|
||||
if errno == 0 {
|
||||
if int(n) < virtio.Size {
|
||||
return 0, fmt.Errorf("tun read shorter than virtio_net_hdr: %d bytes", n)
|
||||
return 0, io.ErrShortWrite
|
||||
}
|
||||
return int(n) - virtio.Size, nil
|
||||
}
|
||||
@@ -167,30 +170,29 @@ func (r *Offload) readPacket(block bool) (int, error) {
|
||||
}
|
||||
}
|
||||
|
||||
// Read returns one or more packets from the tun.
|
||||
// Each Packet either carries a single ready-to-use IP datagram (GSO zero) or a TSO/USO superpacket plus the GSOInfo a caller needs to segment it (see SegmentSuperpacket).
|
||||
// The first read blocks via poll; once the fd is known readable we drain additional packets non-blocking until:
|
||||
// - the kernel queue is empty (EAGAIN)
|
||||
// - we've collected tunDrainCap packets,
|
||||
// - or we're out of rxBuf headroom.
|
||||
//
|
||||
// This amortizes the poll wake over bursts of small packets (e.g. TCP ACKs).
|
||||
// Packet.Bytes slices point into the Offload's internal buffer and are only valid until the next Read or Close on this Queue.
|
||||
// Read returns one or more packets from the tun. Each Packet either
|
||||
// carries a single ready-to-use IP datagram (GSO zero) or a TSO/USO
|
||||
// superpacket plus the GSOInfo a caller needs to segment it (see
|
||||
// SegmentSuperpacket). The first read blocks via poll; once the fd is
|
||||
// known readable we drain additional packets non-blocking until the
|
||||
// kernel queue is empty (EAGAIN), we've collected tunDrainCap packets,
|
||||
// or we're out of rxBuf headroom. This amortizes the poll wake over
|
||||
// bursts of small packets (e.g. TCP ACKs). Packet.Bytes slices point
|
||||
// into the Offload's internal buffer and are only valid until the next
|
||||
// Read or Close on this Queue.
|
||||
func (r *Offload) Read() ([]Packet, error) {
|
||||
r.pending = r.pending[:0]
|
||||
r.rxOff = 0
|
||||
|
||||
// Initial (blocking) read.
|
||||
// Retry on decode errors so a single bad packet does not stall the reader.
|
||||
// Initial (blocking) read. Retry on decode errors so a single bad
|
||||
// packet does not stall the reader.
|
||||
for {
|
||||
n, err := r.readPacket(true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := r.decodeRead(n); err != nil {
|
||||
// Drop and read again. A bad packet should not kill the reader,
|
||||
// but a systematic decode failure must not be invisible either.
|
||||
r.logDroppedRead(err)
|
||||
// Drop and read again — a bad packet should not kill the reader.
|
||||
continue
|
||||
}
|
||||
break
|
||||
@@ -212,7 +214,6 @@ func (r *Offload) Read() ([]Packet, error) {
|
||||
if err := r.decodeRead(n); err != nil {
|
||||
// Drop this packet and stop the drain; we'd rather hand off
|
||||
// what we have than keep spinning here.
|
||||
r.logDroppedRead(err)
|
||||
break
|
||||
}
|
||||
}
|
||||
@@ -220,20 +221,13 @@ func (r *Offload) Read() ([]Packet, error) {
|
||||
return r.pending, nil
|
||||
}
|
||||
|
||||
// logDroppedRead reports a tun packet dropped for a bad/unsupported virtio
|
||||
// header. Debug-gated so the happy path never pays for attribute assembly.
|
||||
func (r *Offload) logDroppedRead(err error) {
|
||||
if r.l != nil && r.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
r.l.Debug("dropping tun packet with bad virtio header", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// decodeRead processes the packet sitting in rxBuf at rxOff (length pktLen).
|
||||
// The bytes stay in rxBuf:
|
||||
// - for GSO_NONE we slice them as a regular IP datagram (running finishChecksum if NEEDS_CSUM is set);
|
||||
// - for TSO/USO superpackets we attach the corrected GSO metadata, so the caller can segment lazily at encrypt time.
|
||||
//
|
||||
// rxOff advances by pktLen on success
|
||||
// decodeRead processes the packet sitting in rxBuf at rxOff (length
|
||||
// pktLen). The bytes stay in rxBuf — for GSO_NONE we slice them as a
|
||||
// regular IP datagram (running finishChecksum if NEEDS_CSUM is set);
|
||||
// for TSO/USO superpackets we attach the corrected GSO metadata so the
|
||||
// caller can segment lazily at encrypt time. rxOff advances past the
|
||||
// kernel-supplied body and nothing else, since segmentation no longer
|
||||
// writes back into rxBuf.
|
||||
func (r *Offload) decodeRead(pktLen int) error {
|
||||
if pktLen <= 0 {
|
||||
return fmt.Errorf("short tun read: %d", pktLen)
|
||||
@@ -243,7 +237,7 @@ func (r *Offload) decodeRead(pktLen int) error {
|
||||
|
||||
body := r.rxBuf[r.rxOff : r.rxOff+pktLen]
|
||||
|
||||
if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
|
||||
if err := virtio.FinishChecksum(body, hdr); err != nil {
|
||||
return err
|
||||
@@ -254,13 +248,17 @@ func (r *Offload) decodeRead(pktLen int) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// GSO superpacket: validate, fix the kernel-supplied HdrLen on the
|
||||
// FORWARD path (CorrectHdrLen), pick the L4 protocol, and attach
|
||||
// the metadata. The bytes stay in rxBuf untouched, segmentation
|
||||
// happens in SegmentSuperpacket at encrypt time.
|
||||
if err := virtio.CheckValid(body, hdr); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := virtio.CorrectHdrLen(body, &hdr); err != nil {
|
||||
return err
|
||||
}
|
||||
proto, err := protoFromGSOType(hdr.GSOType())
|
||||
proto, err := protoFromGSOType(hdr.GSOType)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -278,16 +276,22 @@ func (r *Offload) decodeRead(pktLen int) error {
|
||||
}
|
||||
|
||||
func (r *Offload) Write(buf []byte) (int, error) {
|
||||
if len(buf) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
iovs := [2]unix.Iovec{
|
||||
{Base: &validVnetHdr[0]},
|
||||
{Base: &buf[0]},
|
||||
}
|
||||
iovs[0].SetLen(virtio.Size)
|
||||
iovs[1].SetLen(len(buf))
|
||||
return r.rawWrite(unsafe.Slice(&iovs[0], 2))
|
||||
return r.writeWithScratch(buf, &iovs)
|
||||
}
|
||||
|
||||
func (r *Offload) writeWithScratch(buf []byte, iovs *[2]unix.Iovec) (int, error) {
|
||||
if len(buf) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
iovs[1].Base = &buf[0]
|
||||
iovs[1].SetLen(len(buf))
|
||||
return r.rawWrite(unsafe.Slice(&iovs[0], len(iovs)))
|
||||
}
|
||||
|
||||
func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
|
||||
@@ -317,34 +321,57 @@ func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
|
||||
|
||||
// Capabilities reports the offload features negotiated for this Queue. TSO
|
||||
// is always true for Offload (we only construct it on IFF_VNET_HDR FDs);
|
||||
// USO is true only when the kernel agreed to TUN_F_USO4|6 at open time (Linux ≥ 6.2).
|
||||
// USO is true only when the kernel agreed to TUN_F_USO4|6 at open time
|
||||
// (Linux ≥ 6.2).
|
||||
func (r *Offload) Capabilities() Capabilities {
|
||||
return Capabilities{TSO: true, USO: r.usoEnabled}
|
||||
}
|
||||
|
||||
func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error {
|
||||
if len(pays) == 0 {
|
||||
// There are no payload fragments. There is nothing to send.
|
||||
if len(hdr) == 0 || len(pays) == 0 || len(transportHdr) == 0 {
|
||||
return nil
|
||||
}
|
||||
var csumOff uint16 // csumOff is the offset of the L4 checksum field in transportHdr
|
||||
// L4 checksum offset inside transportHdr: TCP=16 (the `check` field after
|
||||
// seq/ack/dataoff/flags/window), UDP=6 (after sport/dport/length).
|
||||
var csumOff uint16
|
||||
switch proto {
|
||||
case GSOProtoUDP:
|
||||
csumOff = 6
|
||||
case GSOProtoTCP:
|
||||
csumOff = 16
|
||||
default:
|
||||
return fmt.Errorf("unknown GSO proto: %d", proto)
|
||||
csumOff = 16
|
||||
}
|
||||
// Incorrect geometry must cause an error, not a silent drop.
|
||||
// No sane packet should ever make it inside this branch.
|
||||
if len(hdr) == 0 || len(transportHdr) < int(csumOff)+2 {
|
||||
return fmt.Errorf("tio: WriteGSO header too short: ip=%d transport=%d (csum field at %d)", len(hdr), len(transportHdr), csumOff)
|
||||
vhdr := virtio.Hdr{
|
||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||
HdrLen: uint16(len(hdr) + len(transportHdr)),
|
||||
GSOSize: uint16(len(pays[0])),
|
||||
CsumStart: uint16(len(hdr)),
|
||||
CsumOffset: csumOff,
|
||||
}
|
||||
// Make the iovec array: [virtio_hdr, hdr, transportHdr, pays...].
|
||||
// The constructor attaches r.gsoIovs[0] to gsoHdrBuf. That entry does not change.
|
||||
if len(pays) > 1 {
|
||||
ipVer := hdr[0] >> 4
|
||||
switch {
|
||||
case proto == GSOProtoUDP && (ipVer == 4 || ipVer == 6):
|
||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_UDP_L4
|
||||
case ipVer == 6:
|
||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV6
|
||||
case ipVer == 4:
|
||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV4
|
||||
default:
|
||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
||||
vhdr.GSOSize = 0
|
||||
}
|
||||
} else {
|
||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
||||
vhdr.GSOSize = 0
|
||||
}
|
||||
vhdr.Encode(r.gsoHdrBuf[:])
|
||||
|
||||
// Build the iovec array: [virtio_hdr, hdr, transportHdr, pays...]. r.gsoIovs[0] is
|
||||
// wired to gsoHdrBuf at construction and never changes.
|
||||
need := 3 + len(pays)
|
||||
if need > cap(r.gsoIovs) {
|
||||
slog.Default().Warn("tio: WriteGSO iovec budget exceeded; dropping superpacket",
|
||||
"need", need, "cap", cap(r.gsoIovs), "segments", len(pays))
|
||||
return fmt.Errorf("tio: WriteGSO needs %d iovecs but cap is %d", need, cap(r.gsoIovs))
|
||||
}
|
||||
r.gsoIovs = r.gsoIovs[:need]
|
||||
@@ -352,52 +379,22 @@ func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto
|
||||
r.gsoIovs[1].SetLen(len(hdr))
|
||||
r.gsoIovs[2].Base = &transportHdr[0]
|
||||
r.gsoIovs[2].SetLen(len(transportHdr))
|
||||
|
||||
segSize := len(pays[0])
|
||||
total := len(hdr) + len(transportHdr)
|
||||
for i, p := range pays {
|
||||
// Defense in depth: an empty payload fragment can't be a valid GSO
|
||||
// segment and &p[0] would panic on it. Callers route zero-length
|
||||
// datagrams through the plain path (see UDPCoalescer.commitParsed), so
|
||||
// this should never fire, but skip empties rather than index into one.
|
||||
// `n` tracks where the next payload iovec lands, since skips make it
|
||||
// drift from 3+i.
|
||||
n := 3
|
||||
for _, p := range pays {
|
||||
if len(p) == 0 {
|
||||
// The coalescers route zero-payload packets down the non-GSO path,
|
||||
// so an empty fragment means the caller's accounting is broken.
|
||||
return fmt.Errorf("tio: WriteGSO empty payload fragment %d of %d", i, len(pays))
|
||||
} else if len(p) > segSize || (len(p) < segSize && i != len(pays)-1) {
|
||||
// all segments must be the same size, except for the last one
|
||||
return fmt.Errorf("tio: WriteGSO fragment %d is %dB, want %dB segments (only the last may be shorter)", i, len(p), segSize)
|
||||
continue
|
||||
}
|
||||
total += len(p)
|
||||
r.gsoIovs[3+i].Base = &p[0]
|
||||
r.gsoIovs[3+i].SetLen(len(p))
|
||||
r.gsoIovs[n].Base = &p[0]
|
||||
r.gsoIovs[n].SetLen(len(p))
|
||||
n++
|
||||
}
|
||||
// This check keeps `total` in the uint16 range. Anything larger would wrap around and cause trouble.
|
||||
if total > maxSuperpacketLen {
|
||||
return fmt.Errorf("tio: WriteGSO superpacket %dB exceeds %d", total, maxSuperpacketLen)
|
||||
}
|
||||
|
||||
// A single segment ships as a plain checksummed packet (GSO_NONE, size 0).
|
||||
// Multiple segments carry the real GSO type and segSize, which the loop
|
||||
// above verified is the size of every fragment except possibly the last.
|
||||
gsoType := uint8(unix.VIRTIO_NET_HDR_GSO_NONE)
|
||||
if len(pays) > 1 {
|
||||
gsoType = gsoTypeFromProto(proto, hdr[0]>>4)
|
||||
if gsoType == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||
// gsoTypeFromProto only yields GSO_NONE for a bogus IP version nibble.
|
||||
// A multi-segment superpacket must carry a real GSO type, or the kernel would deliver it as a single jumbo packet.
|
||||
return fmt.Errorf("tio: WriteGSO IP version %d is not GSO-capable", hdr[0]>>4)
|
||||
}
|
||||
}
|
||||
var gsoSize uint16
|
||||
if gsoType != unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||
gsoSize = uint16(segSize)
|
||||
}
|
||||
virtio.EncodeHeader(
|
||||
r.gsoHdrBuf[:],
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
gsoType, /*gsoType*/
|
||||
uint16(len(hdr)+len(transportHdr)), /*hdrLen*/
|
||||
gsoSize, /*gsoSize*/
|
||||
uint16(len(hdr)), /*csumStart*/
|
||||
csumOff, /*csumOffset*/
|
||||
)
|
||||
r.gsoIovs = r.gsoIovs[:n]
|
||||
|
||||
_, err := r.rawWrite(r.gsoIovs)
|
||||
return err
|
||||
@@ -408,10 +405,10 @@ func (r *Offload) Close() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// shutdownFd is owned by the container, so we should not close it
|
||||
// Close the underlying fd but do NOT null r.fd: a reader may still be loading it in readPacket, and mutating the field would race that load.
|
||||
// That reader gets EBADF -> os.ErrClosed on its next syscall. A reader already parked in
|
||||
// poll is NOT woken by this close; only the QueueSet's shutdown eventfd wake does that (see Queue.Close docs).
|
||||
// closed.Swap already guarantees we only close once.
|
||||
//shutdownFd is owned by the container, so we should not close it
|
||||
// Close the underlying fd but do NOT null r.fd: a reader may still be
|
||||
// loading it in readOne, and mutating the field would race that load.
|
||||
// It gets EBADF -> os.ErrClosed (or wakes via the shutdown eventfd's
|
||||
// ppoll first). closed.Swap already guarantees we only close once.
|
||||
return unix.Close(r.fd)
|
||||
}
|
||||
|
||||
@@ -11,6 +11,11 @@ import (
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// Maximum size we accept for a single read from a TUN with IFF_VNET_HDR. A
|
||||
// TSO superpacket can be up to 64KiB of payload plus a single L2/L3/L4 header
|
||||
// prefix plus the virtio header.
|
||||
const tunReadBufSize = 65535
|
||||
|
||||
type Poll struct {
|
||||
fd int
|
||||
shutdownFd int
|
||||
@@ -20,10 +25,10 @@ type Poll struct {
|
||||
batchRet [1]Packet
|
||||
}
|
||||
|
||||
// newPoll wraps an existing tun fd.
|
||||
// On failure it does NOT close fd: the caller owns fd and is the sole closer
|
||||
// (see pollQueueSet.Add callers in overlay/tun_linux.go, which unix.Close on Add error).
|
||||
// This matches the newOffload convention and keeps closes at exactly one on every path.
|
||||
// newPoll wraps an existing tun fd. On failure it does NOT close fd: the
|
||||
// caller owns fd and is the sole closer (see pollQueueSet.Add callers in
|
||||
// overlay/tun_linux.go, which unix.Close on Add error). This matches the
|
||||
// newOffload convention and keeps closes at exactly one on every path.
|
||||
func newPoll(fd int, shutdownFd int) (*Poll, error) {
|
||||
if err := unix.SetNonblock(fd, true); err != nil {
|
||||
return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err)
|
||||
@@ -32,7 +37,7 @@ func newPoll(fd int, shutdownFd int) (*Poll, error) {
|
||||
out := &Poll{
|
||||
fd: fd,
|
||||
shutdownFd: shutdownFd,
|
||||
readBuf: make([]byte, 65535), // largest possible size Linux permits
|
||||
readBuf: make([]byte, tunReadBufSize),
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
@@ -47,11 +52,6 @@ func (t *Poll) blockOnWrite() error {
|
||||
return blockOn(int32(t.fd), int32(t.shutdownFd), unix.POLLOUT)
|
||||
}
|
||||
|
||||
// TODO: port Offload's post-wake drain loop here so one poll wake amortizes
|
||||
// over a burst (up to tunDrainCap packets) instead of paying a syscall and a
|
||||
// wake per packet. Hosts on the TUNSETOFFLOAD-failure fallback or a tun.fd
|
||||
// config currently lose that batching. blockOn and the EAGAIN plumbing are
|
||||
// already shared; kept one-packet-per-Read for now to preserve behavior.
|
||||
func (t *Poll) Read() ([]Packet, error) {
|
||||
n, err := t.readOne(t.readBuf)
|
||||
if err != nil {
|
||||
@@ -108,11 +108,10 @@ func (t *Poll) Close() error {
|
||||
if t.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
|
||||
// shutdownFd is owned by the container, so we should not close it
|
||||
// Close the underlying fd but do NOT null t.fd: a reader may still be loading it in readOne, and mutating the field would race that load.
|
||||
// That reader gets EBADF -> os.ErrClosed on its next syscall. A reader already parked in
|
||||
// poll is NOT woken by this close; only the QueueSet's shutdown eventfd wake does that (see Queue.Close docs).
|
||||
// closed.Swap already guarantees we only close once.
|
||||
//shutdownFd is owned by the container, so we should not close it
|
||||
// Close the underlying fd but do NOT null t.fd: a reader may still be
|
||||
// loading it in readOne, and mutating the field would race that load.
|
||||
// It gets EBADF -> os.ErrClosed (or wakes via the shutdown eventfd's
|
||||
// ppoll first). closed.Swap already guarantees we only close once.
|
||||
return unix.Close(t.fd)
|
||||
}
|
||||
|
||||
@@ -5,7 +5,6 @@ package tio
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"log/slog"
|
||||
"os"
|
||||
"sync"
|
||||
"testing"
|
||||
@@ -211,7 +210,7 @@ func TestPollQueueSet_Close_ClosesShutdownFd(t *testing.T) {
|
||||
// TestOffloadQueueSet_Close_ClosesShutdownFd mirrors the poll regression test
|
||||
// for the GSO/offload queueset.
|
||||
func TestOffloadQueueSet_Close_ClosesShutdownFd(t *testing.T) {
|
||||
qs, err := NewOffloadQueueSet(false, slog.New(slog.DiscardHandler))
|
||||
qs, err := NewOffloadQueueSet(false)
|
||||
require.NoError(t, err)
|
||||
c, ok := qs.(*offloadQueueSet)
|
||||
require.True(t, ok)
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
//go:build linux && !android && !e2e_testing
|
||||
// +build linux,!android,!e2e_testing
|
||||
|
||||
package tio
|
||||
|
||||
@@ -11,11 +11,11 @@ import (
|
||||
"github.com/slackhq/nebula/overlay/tio/virtio"
|
||||
)
|
||||
|
||||
// protoFromGSOType maps a virtio_net_hdr gsoType to the GSOProto value the
|
||||
// protoFromGSOType maps a virtio_net_hdr GSOType to the GSOProto value the
|
||||
// segment-time helpers use. Returns an error for GSO_NONE or any unknown
|
||||
// value. The caller should only invoke this on a confirmed superpacket.
|
||||
// value — the caller should only invoke this on a confirmed superpacket.
|
||||
func protoFromGSOType(t uint8) (GSOProto, error) {
|
||||
switch t &^ unix.VIRTIO_NET_HDR_GSO_ECN {
|
||||
switch t {
|
||||
case unix.VIRTIO_NET_HDR_GSO_TCPV4, unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
||||
return GSOProtoTCP, nil
|
||||
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
|
||||
@@ -25,26 +25,17 @@ func protoFromGSOType(t uint8) (GSOProto, error) {
|
||||
}
|
||||
}
|
||||
|
||||
// gsoTypeFromProto is the reverse of protoFromGSOType
|
||||
func gsoTypeFromProto(proto GSOProto, ipVer uint8) uint8 {
|
||||
switch {
|
||||
case proto == GSOProtoUDP && (ipVer == 4 || ipVer == 6):
|
||||
return unix.VIRTIO_NET_HDR_GSO_UDP_L4
|
||||
case ipVer == 6:
|
||||
return unix.VIRTIO_NET_HDR_GSO_TCPV6
|
||||
case ipVer == 4:
|
||||
return unix.VIRTIO_NET_HDR_GSO_TCPV4
|
||||
default:
|
||||
return unix.VIRTIO_NET_HDR_GSO_NONE
|
||||
}
|
||||
}
|
||||
|
||||
// SegmentSuperpacket invokes fn once per segment of pkt.
|
||||
// For non-GSO pkts fn is called once with pkt.Bytes.
|
||||
// For GSO/USO superpackets, fn is called once per segment with a slice of pkt.Bytes holding that segment's plaintext
|
||||
// (a freshly-patched L3+L4 header sliced in front of the original payload chunk).
|
||||
// This slicing is destructive: pkt is consumed by this call.
|
||||
// Aborts and returns the first error from fn or from per-segment construction.
|
||||
// SegmentSuperpacket invokes fn once per segment of pkt. For non-GSO pkts
|
||||
// fn is called once with pkt.Bytes (no segmentation, no copy). For GSO/USO
|
||||
// superpackets fn is called once per segment with a slice of pkt.Bytes
|
||||
// holding that segment's plaintext (a freshly-patched L3+L4 header sliced
|
||||
// in front of the original payload chunk). The slide is destructive: pkt is
|
||||
// consumed by this call and its bytes are in an undefined state when
|
||||
// SegmentSuperpacket returns. Callers must not retain pkt or any earlier
|
||||
// seg slice past fn's return for that segment. The scratch parameter is
|
||||
// unused on the destructive path and kept only for cross-platform
|
||||
// signature compatibility. Aborts and returns the first error from fn or
|
||||
// from per-segment construction.
|
||||
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
||||
if !pkt.GSO.IsSuperpacket() {
|
||||
return fn(pkt.Bytes)
|
||||
|
||||
@@ -19,36 +19,6 @@ import (
|
||||
// worst-case 64 KiB superpacket plus replicated per-segment headers).
|
||||
const testSegScratchSize = 192 * 1024
|
||||
|
||||
// TestProtoFromGSOTypeMasksECN guards the CWR-superpacket drop bug: the
|
||||
// kernel qualifies a TSO superpacket whose TCP header carries CWR with
|
||||
// VIRTIO_NET_HDR_GSO_ECN (we negotiate TUN_F_TSO_ECN, so it WILL send
|
||||
// them once ECN feedback flows), and the decoder must mask that bit
|
||||
// rather than reject the packet as an unknown type.
|
||||
func TestProtoFromGSOTypeMasksECN(t *testing.T) {
|
||||
cases := []struct {
|
||||
typ uint8
|
||||
want GSOProto
|
||||
}{
|
||||
{unix.VIRTIO_NET_HDR_GSO_TCPV4, GSOProtoTCP},
|
||||
{unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, GSOProtoTCP},
|
||||
{unix.VIRTIO_NET_HDR_GSO_TCPV6, GSOProtoTCP},
|
||||
{unix.VIRTIO_NET_HDR_GSO_TCPV6 | unix.VIRTIO_NET_HDR_GSO_ECN, GSOProtoTCP},
|
||||
{unix.VIRTIO_NET_HDR_GSO_UDP_L4, GSOProtoUDP},
|
||||
}
|
||||
for _, c := range cases {
|
||||
got, err := protoFromGSOType(c.typ)
|
||||
if err != nil || got != c.want {
|
||||
t.Errorf("protoFromGSOType(%#x) = (%v, %v), want (%v, nil)", c.typ, got, err, c.want)
|
||||
}
|
||||
}
|
||||
if _, err := protoFromGSOType(unix.VIRTIO_NET_HDR_GSO_NONE); err == nil {
|
||||
t.Error("GSO_NONE must still be rejected")
|
||||
}
|
||||
if _, err := protoFromGSOType(unix.VIRTIO_NET_HDR_GSO_ECN); err == nil {
|
||||
t.Error("a bare ECN bit with no base type must still be rejected")
|
||||
}
|
||||
}
|
||||
|
||||
// verifyChecksum confirms that the one's-complement sum across `b`, seeded
|
||||
// with a folded pseudo-header sum, equals all-ones (valid).
|
||||
func verifyChecksum(b []byte, pseudo uint16) bool {
|
||||
@@ -63,7 +33,7 @@ func verifyChecksum(b []byte, pseudo uint16) bool {
|
||||
// returns. Tests pre-set hdr.HdrLen correctly, so correctHdrLen is not
|
||||
// invoked here.
|
||||
func segmentForTest(pkt []byte, hdr virtio.Hdr, out *[][]byte, scratch []byte) error {
|
||||
if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||
cp := append([]byte(nil), pkt...)
|
||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
|
||||
if err := virtio.FinishChecksum(cp, hdr); err != nil {
|
||||
@@ -73,7 +43,7 @@ func segmentForTest(pkt []byte, hdr virtio.Hdr, out *[][]byte, scratch []byte) e
|
||||
*out = append(*out, cp)
|
||||
return nil
|
||||
}
|
||||
proto, err := protoFromGSOType(hdr.GSOType())
|
||||
proto, err := protoFromGSOType(hdr.GSOType)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -140,14 +110,15 @@ func buildTSOv4(t *testing.T, payLen, mss int) ([]byte, virtio.Hdr) {
|
||||
for i := 0; i < payLen; i++ {
|
||||
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
|
||||
}
|
||||
return pkt, virtio.NewHeader(
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
unix.VIRTIO_NET_HDR_GSO_TCPV4, /*gsoType*/
|
||||
uint16(ipLen+tcpLen), /*hdrLen*/
|
||||
uint16(mss), /*gsoSize*/
|
||||
uint16(ipLen), /*csumStart*/
|
||||
16, /*csumOffset*/
|
||||
)
|
||||
|
||||
return pkt, virtio.Hdr{
|
||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4,
|
||||
HdrLen: uint16(ipLen + tcpLen),
|
||||
GSOSize: uint16(mss),
|
||||
CsumStart: uint16(ipLen),
|
||||
CsumOffset: 16,
|
||||
}
|
||||
}
|
||||
|
||||
func TestSegmentTCPv4(t *testing.T) {
|
||||
@@ -261,14 +232,14 @@ func TestSegmentTCPv6(t *testing.T) {
|
||||
pkt[ipLen+tcpLen+i] = byte(i)
|
||||
}
|
||||
|
||||
hdr := virtio.NewHeader(
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
unix.VIRTIO_NET_HDR_GSO_TCPV6, /*gsoType*/
|
||||
uint16(ipLen+tcpLen), /*hdrLen*/
|
||||
uint16(mss), /*gsoSize*/
|
||||
uint16(ipLen), /*csumStart*/
|
||||
16, /*csumOffset*/
|
||||
)
|
||||
hdr := virtio.Hdr{
|
||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV6,
|
||||
HdrLen: uint16(ipLen + tcpLen),
|
||||
GSOSize: uint16(mss),
|
||||
CsumStart: uint16(ipLen),
|
||||
CsumOffset: 16,
|
||||
}
|
||||
|
||||
scratch := make([]byte, testSegScratchSize)
|
||||
var out [][]byte
|
||||
@@ -310,7 +281,7 @@ func TestSegmentTCPv6(t *testing.T) {
|
||||
|
||||
func TestSegmentGSONonePassesThrough(t *testing.T) {
|
||||
pkt, hdr := buildTSOv4(t, 100, 100)
|
||||
hdr.SetGSOType(unix.VIRTIO_NET_HDR_GSO_NONE)
|
||||
hdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
||||
hdr.Flags = 0 // no NEEDS_CSUM, leave packet untouched
|
||||
|
||||
scratch := make([]byte, testSegScratchSize)
|
||||
@@ -329,7 +300,7 @@ func TestSegmentGSONonePassesThrough(t *testing.T) {
|
||||
// TestSegmentRejectsLegacyUDPGSO ensures the legacy GSO_UDP (UFO) marker is
|
||||
// still rejected; only modern GSO_UDP_L4 (USO) is supported.
|
||||
func TestSegmentRejectsLegacyUDPGSO(t *testing.T) {
|
||||
hdr := virtio.NewHeader(0, unix.VIRTIO_NET_HDR_GSO_UDP, 0, 0, 0, 0)
|
||||
hdr := virtio.Hdr{GSOType: unix.VIRTIO_NET_HDR_GSO_UDP}
|
||||
var out [][]byte
|
||||
if err := segmentForTest(nil, hdr, &out, nil); err == nil {
|
||||
t.Fatalf("expected rejection for legacy UDP GSO")
|
||||
@@ -353,26 +324,22 @@ func buildUSOv4(t *testing.T, payLen, gsoSize int) ([]byte, virtio.Hdr) {
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||
|
||||
// UDP header. The kernel hands us a USO superpacket whose length field
|
||||
// covers the WHOLE superpacket; the segmenter overwrites it per segment.
|
||||
// Populating it here matters: leaving it zero makes the base-checksum path
|
||||
// that must exclude it untestable, since excluding zero is a no-op.
|
||||
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
|
||||
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
|
||||
binary.BigEndian.PutUint16(pkt[24:26], uint16(udpLen+payLen)) // superpacket length
|
||||
// UDP header (length + checksum filled in per segment by segmentUDPYield)
|
||||
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
|
||||
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
|
||||
|
||||
for i := 0; i < payLen; i++ {
|
||||
pkt[ipLen+udpLen+i] = byte(i & 0xff)
|
||||
}
|
||||
|
||||
return pkt, virtio.NewHeader(
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
|
||||
uint16(ipLen+udpLen), /*hdrLen*/
|
||||
uint16(gsoSize), /*gsoSize*/
|
||||
uint16(ipLen), /*csumStart*/
|
||||
6, /*csumOffset*/
|
||||
)
|
||||
return pkt, virtio.Hdr{
|
||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
|
||||
HdrLen: uint16(ipLen + udpLen),
|
||||
GSOSize: uint16(gsoSize),
|
||||
CsumStart: uint16(ipLen),
|
||||
CsumOffset: 6,
|
||||
}
|
||||
}
|
||||
|
||||
func TestSegmentUDPv4(t *testing.T) {
|
||||
@@ -397,12 +364,11 @@ func TestSegmentUDPv4(t *testing.T) {
|
||||
if totalLen != uint16(28+gso) {
|
||||
t.Errorf("seg %d: total_len=%d want %d", i, totalLen, 28+gso)
|
||||
}
|
||||
// Software UDP GSO bumps the IPv4 ID per segment exactly like TSO
|
||||
// (inet_gso_segment's fixed-ID case is TCP-only); wireguard-go's
|
||||
// gsoSplit increments unconditionally too.
|
||||
// kernel UDP-GSO does NOT bump the IPv4 ID across segments; every
|
||||
// segment carries the same ID as the seed.
|
||||
id := binary.BigEndian.Uint16(seg[4:6])
|
||||
if id != 0x4242+uint16(i) {
|
||||
t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242+uint16(i))
|
||||
if id != 0x4242 {
|
||||
t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242)
|
||||
}
|
||||
udpLen := binary.BigEndian.Uint16(seg[24:26])
|
||||
if udpLen != uint16(8+gso) {
|
||||
@@ -470,21 +436,19 @@ func TestSegmentUDPv6(t *testing.T) {
|
||||
|
||||
binary.BigEndian.PutUint16(pkt[40:42], 12345)
|
||||
binary.BigEndian.PutUint16(pkt[42:44], 53)
|
||||
// Superpacket-wide length, as the kernel supplies it; see buildUSOv4.
|
||||
binary.BigEndian.PutUint16(pkt[44:46], uint16(udpLen+payLen))
|
||||
|
||||
for i := 0; i < payLen; i++ {
|
||||
pkt[ipLen+udpLen+i] = byte(i)
|
||||
}
|
||||
|
||||
hdr := virtio.NewHeader(
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
|
||||
uint16(ipLen+udpLen), /*hdrLen*/
|
||||
uint16(gso), /*gsoSize*/
|
||||
uint16(ipLen), /*csumStart*/
|
||||
6, /*csumOffset*/
|
||||
)
|
||||
hdr := virtio.Hdr{
|
||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
|
||||
HdrLen: uint16(ipLen + udpLen),
|
||||
GSOSize: uint16(gso),
|
||||
CsumStart: uint16(ipLen),
|
||||
CsumOffset: 6,
|
||||
}
|
||||
|
||||
scratch := make([]byte, testSegScratchSize)
|
||||
var out [][]byte
|
||||
@@ -616,14 +580,14 @@ func BenchmarkSegmentTCPv4(b *testing.B) {
|
||||
for i := 0; i < sz.payLen; i++ {
|
||||
pkt[ipLen+tcpLen+i] = byte(i)
|
||||
}
|
||||
hdr := virtio.NewHeader(
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
unix.VIRTIO_NET_HDR_GSO_TCPV4, /*gsoType*/
|
||||
uint16(ipLen+tcpLen), /*hdrLen*/
|
||||
uint16(sz.mss), /*gsoSize*/
|
||||
uint16(ipLen), /*csumStart*/
|
||||
16, /*csumOffset*/
|
||||
)
|
||||
hdr := virtio.Hdr{
|
||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4,
|
||||
HdrLen: uint16(ipLen + tcpLen),
|
||||
GSOSize: uint16(sz.mss),
|
||||
CsumStart: uint16(ipLen),
|
||||
CsumOffset: 16,
|
||||
}
|
||||
|
||||
scratch := make([]byte, testSegScratchSize)
|
||||
out := make([][]byte, 0, 64)
|
||||
@@ -676,72 +640,35 @@ func TestTunFileWriteVnetHdrNoAlloc(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestSegmentSuperpacketNoAlloc pins the segmenters' zero-allocation
|
||||
// contract. Both SegmentTCP and SegmentUDP derive their per-superpacket
|
||||
// constants into fixed-size arrays (tmp/ipTmp/savedHdr) that must stay on
|
||||
// the stack, and both take a yield closure that must not escape. Any of
|
||||
// those escaping turns one allocation into one-per-superpacket on the
|
||||
// hottest path in the reader, which BenchmarkSegmentSuperpacketAllocsTSO
|
||||
// reports but nothing fails on. This does.
|
||||
//
|
||||
// The yield closure here only touches captured scalars: appending segments
|
||||
// to a slice would allocate in the test itself and mask the measurement.
|
||||
func TestSegmentSuperpacketNoAlloc(t *testing.T) {
|
||||
const mss = 1400
|
||||
const numSeg = 8
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
build func() ([]byte, virtio.Hdr)
|
||||
}{
|
||||
{"tso-v4", func() ([]byte, virtio.Hdr) { return buildTSOv4(t, mss*numSeg, mss) }},
|
||||
{"uso-v4", func() ([]byte, virtio.Hdr) { return buildUSOv4(t, mss*numSeg, mss) }},
|
||||
// TestWriteGSOSkipsEmptyPayloads is the defense-in-depth guard for the
|
||||
// zero-length UDP DoS: a payload fragment of length zero would make &p[0]
|
||||
// panic (index-out-of-range) when building the iovec array. WriteGSO must
|
||||
// skip empties instead. We write to /dev/null so the writev always succeeds
|
||||
// synchronously; the point is simply that neither call panics.
|
||||
func TestWriteGSOSkipsEmptyPayloads(t *testing.T) {
|
||||
fd, err := unix.Open("/dev/null", os.O_WRONLY, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("open /dev/null: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = unix.Close(fd) })
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
master, hdr := tc.build()
|
||||
proto, err := protoFromGSOType(hdr.GSOType())
|
||||
if err != nil {
|
||||
t.Fatalf("protoFromGSOType: %v", err)
|
||||
}
|
||||
work := make([]byte, len(master))
|
||||
p := Packet{Bytes: work, GSO: GSOInfo{
|
||||
Size: hdr.GSOSize,
|
||||
HdrLen: hdr.HdrLen,
|
||||
CsumStart: hdr.CsumStart,
|
||||
Proto: proto,
|
||||
}}
|
||||
o := &Offload{fd: fd, gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs)}
|
||||
o.gsoIovs[0].Base = &o.gsoHdrBuf[0]
|
||||
o.gsoIovs[0].SetLen(virtio.Size)
|
||||
|
||||
// Segmentation consumes its input destructively, so restore from
|
||||
// the master copy each run; copy(2) into an existing slice does
|
||||
// not allocate. seen/bytes keep the closure from being optimized
|
||||
// away and double as a sanity check that work actually happened.
|
||||
var seen, bytes int
|
||||
run := func() {
|
||||
copy(work, master)
|
||||
seen, bytes = 0, 0
|
||||
if err := SegmentSuperpacket(p, func(seg []byte) error {
|
||||
seen++
|
||||
bytes += len(seg)
|
||||
return nil
|
||||
}); err != nil {
|
||||
t.Fatalf("SegmentSuperpacket: %v", err)
|
||||
}
|
||||
}
|
||||
ipHdr := make([]byte, 20)
|
||||
ipHdr[0] = 0x45 // IPv4, IHL 5
|
||||
udpHdr := make([]byte, 8)
|
||||
|
||||
run() // warm up: absorb any one-time allocation elsewhere
|
||||
if seen != numSeg {
|
||||
t.Fatalf("yielded %d segments, want %d", seen, numSeg)
|
||||
}
|
||||
|
||||
if allocs := testing.AllocsPerRun(200, run); allocs != 0 {
|
||||
t.Fatalf("SegmentSuperpacket allocated %.1f times per call, want 0", allocs)
|
||||
}
|
||||
if seen != numSeg || bytes == 0 {
|
||||
t.Fatalf("post-measure sanity: seen=%d bytes=%d", seen, bytes)
|
||||
}
|
||||
})
|
||||
// Sole payload empty: exercises the all-empty skip (n stays at 3).
|
||||
if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{{}}, GSOProtoUDP); err != nil {
|
||||
t.Fatalf("WriteGSO with a single empty payload: %v", err)
|
||||
}
|
||||
// Empty mixed with a real fragment: exercises the index-drift skip so a
|
||||
// later non-empty payload still lands in the right iovec slot.
|
||||
real := make([]byte, 1200)
|
||||
if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{real, {}}, GSOProtoUDP); err != nil {
|
||||
t.Fatalf("WriteGSO with a trailing empty payload: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -794,11 +721,9 @@ func TestDecodeReadFitsMaxTSOAtDrainThreshold(t *testing.T) {
|
||||
const ipv6HdrLen = 40
|
||||
const tcpHdrLen = 20
|
||||
const headerLen = ipv6HdrLen + tcpHdrLen
|
||||
// Maximum TUN read body at the drain threshold. readv bounds the body
|
||||
// iovec by the space actually left in rxBuf, and the drain gate keeps that
|
||||
// at >= tunRxBufSize, so that is the largest superpacket the kernel can
|
||||
// hand back on the last permitted drain read.
|
||||
pktLen := tunRxBufSize
|
||||
// Maximum TUN read body. The tunReadBufSize cap on readv's body iovec
|
||||
// is what bounds the kernel's superpacket length.
|
||||
pktLen := tunReadBufSize
|
||||
payLen := pktLen - headerLen
|
||||
const targetSegs = 64
|
||||
gsoSize := (payLen + targetSegs - 1) / targetSegs
|
||||
@@ -820,14 +745,14 @@ func TestDecodeReadFitsMaxTSOAtDrainThreshold(t *testing.T) {
|
||||
copy(o.rxBuf[o.rxOff:], pkt)
|
||||
|
||||
// Encode the matching virtio_net_hdr.
|
||||
hdr := virtio.NewHeader(
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
unix.VIRTIO_NET_HDR_GSO_TCPV6, /*gsoType*/
|
||||
uint16(headerLen), /*hdrLen*/
|
||||
uint16(gsoSize), /*gsoSize*/
|
||||
uint16(ipv6HdrLen), /*csumStart*/
|
||||
16, /*csumOffset*/
|
||||
)
|
||||
hdr := virtio.Hdr{
|
||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV6,
|
||||
HdrLen: uint16(headerLen),
|
||||
GSOSize: uint16(gsoSize),
|
||||
CsumStart: uint16(ipv6HdrLen),
|
||||
CsumOffset: 16,
|
||||
}
|
||||
hdr.Encode(o.readVnetScratch[:])
|
||||
|
||||
startRxOff := o.rxOff
|
||||
@@ -899,224 +824,3 @@ func TestDecodeReadFitsMaxTSOAtDrainThreshold(t *testing.T) {
|
||||
t.Fatalf("got %d segments, want %d", gotSegs, wantSegs)
|
||||
}
|
||||
}
|
||||
|
||||
// TestOffloadWriteZeroLength: a zero-length Write must be a no-op, not a
|
||||
// panic. The guard used to live below the &buf[0] that tripped on it.
|
||||
func TestOffloadWriteZeroLength(t *testing.T) {
|
||||
tf := &Offload{fd: -1} // any write reaching the fd would fail loudly
|
||||
for _, buf := range [][]byte{nil, {}} {
|
||||
n, err := tf.Write(buf)
|
||||
if n != 0 || err != nil {
|
||||
t.Errorf("Write(len=0) = (%d, %v), want (0, nil)", n, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteGSOSuperpacketGeometry decodes the vnet header the kernel would see for a multi-segment write:
|
||||
// the GSO type must match the proto and IP version
|
||||
// gso_size must be the per-segment size (the kernel rejects a superpacket with gso_size == 0),
|
||||
// and the csum fields must point at the transport header's checksum slot.
|
||||
// Write through a pipe so the bytes can be read back and decoded.
|
||||
func TestWriteGSOSuperpacketGeometry(t *testing.T) {
|
||||
var pfds [2]int
|
||||
if err := unix.Pipe(pfds[:]); err != nil {
|
||||
t.Fatalf("pipe: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { unix.Close(pfds[0]); unix.Close(pfds[1]) })
|
||||
|
||||
o := &Offload{fd: pfds[1], gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs)}
|
||||
o.gsoIovs[0].Base = &o.gsoHdrBuf[0]
|
||||
o.gsoIovs[0].SetLen(virtio.Size)
|
||||
|
||||
ipHdr := make([]byte, 20)
|
||||
ipHdr[0] = 0x45
|
||||
udpHdr := make([]byte, 8)
|
||||
seg := make([]byte, 1200)
|
||||
|
||||
if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{seg, seg}, GSOProtoUDP); err != nil {
|
||||
t.Fatalf("WriteGSO: %v", err)
|
||||
}
|
||||
|
||||
buf := make([]byte, virtio.Size+len(ipHdr)+len(udpHdr)+2*len(seg)+64)
|
||||
n, err := unix.Read(pfds[0], buf)
|
||||
if err != nil {
|
||||
t.Fatalf("read pipe: %v", err)
|
||||
}
|
||||
var vhdr virtio.Hdr
|
||||
vhdr.Decode(buf[:virtio.Size])
|
||||
if vhdr.GSOType() != unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
|
||||
t.Errorf("gsoType=%d want UDP_L4", vhdr.GSOType())
|
||||
}
|
||||
if vhdr.GSOSize != 1200 {
|
||||
t.Errorf("GSOSize=%d want 1200 (per-segment size from pays[0])", vhdr.GSOSize)
|
||||
}
|
||||
if vhdr.HdrLen != uint16(len(ipHdr)+len(udpHdr)) {
|
||||
t.Errorf("HdrLen=%d want %d", vhdr.HdrLen, len(ipHdr)+len(udpHdr))
|
||||
}
|
||||
if vhdr.CsumStart != uint16(len(ipHdr)) || vhdr.CsumOffset != 6 {
|
||||
t.Errorf("csum start/offset = %d/%d want %d/6", vhdr.CsumStart, vhdr.CsumOffset, len(ipHdr))
|
||||
}
|
||||
if want := virtio.Size + len(ipHdr) + len(udpHdr) + 2*len(seg); n != want {
|
||||
t.Errorf("wrote %d bytes want %d", n, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteGSORejectsBadGeometry pins the length-check contracts
|
||||
func TestWriteGSORejectsBadGeometry(t *testing.T) {
|
||||
fd, err := unix.Open("/dev/null", os.O_WRONLY, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("open /dev/null: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = unix.Close(fd) })
|
||||
|
||||
o := &Offload{fd: fd, gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs)}
|
||||
o.gsoIovs[0].Base = &o.gsoHdrBuf[0]
|
||||
o.gsoIovs[0].SetLen(virtio.Size)
|
||||
|
||||
ipHdr := make([]byte, 20)
|
||||
ipHdr[0] = 0x45
|
||||
udpHdr := make([]byte, 8)
|
||||
tcpHdr := make([]byte, 20)
|
||||
seg := make([]byte, 1200)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
hdr, thdr []byte
|
||||
pays [][]byte
|
||||
proto GSOProto
|
||||
wantErr bool
|
||||
}{
|
||||
{"empty-ip-hdr-with-payload", nil, udpHdr, [][]byte{seg}, GSOProtoUDP, true},
|
||||
{"udp-transport-too-short-for-csum", ipHdr, udpHdr[:6], [][]byte{seg}, GSOProtoUDP, true},
|
||||
{"tcp-transport-too-short-for-csum", ipHdr, tcpHdr[:16], [][]byte{seg}, GSOProtoTCP, true},
|
||||
{"superpacket-over-65535", ipHdr, tcpHdr, [][]byte{make([]byte, 40000), make([]byte, 40000)}, GSOProtoTCP, true},
|
||||
{"sole-payload-empty", ipHdr, udpHdr, [][]byte{{}}, GSOProtoUDP, true},
|
||||
{"leading-empty-fragment", ipHdr, udpHdr, [][]byte{{}, seg, seg}, GSOProtoUDP, true},
|
||||
{"trailing-empty-fragment", ipHdr, tcpHdr, [][]byte{seg, {}}, GSOProtoTCP, true},
|
||||
{"oversize-middle-fragment", ipHdr, udpHdr, [][]byte{seg, make([]byte, 1201), seg}, GSOProtoUDP, true},
|
||||
{"undersize-middle-fragment", ipHdr, udpHdr, [][]byte{seg, make([]byte, 100), seg}, GSOProtoUDP, true},
|
||||
{"oversize-last-fragment", ipHdr, tcpHdr, [][]byte{seg, make([]byte, 1201)}, GSOProtoTCP, true},
|
||||
{"short-last-fragment-ok", ipHdr, udpHdr, [][]byte{seg, seg, make([]byte, 100)}, GSOProtoUDP, false},
|
||||
{"multi-segment-bad-ip-version", []byte{0x05}, udpHdr, [][]byte{seg, seg}, GSOProtoUDP, true},
|
||||
{"single-segment-bad-ip-version-ok", []byte{0x05}, udpHdr, [][]byte{seg}, GSOProtoUDP, false},
|
||||
{"no-pays-noop", ipHdr, udpHdr, nil, GSOProtoUDP, false},
|
||||
{"valid-udp", ipHdr, udpHdr, [][]byte{seg, seg}, GSOProtoUDP, false},
|
||||
{"valid-tcp", ipHdr, tcpHdr, [][]byte{seg, seg}, GSOProtoTCP, false},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
err := o.WriteGSO(tc.hdr, tc.thdr, tc.pays, tc.proto)
|
||||
if tc.wantErr && err == nil {
|
||||
t.Errorf("WriteGSO = nil, want error")
|
||||
}
|
||||
if !tc.wantErr && err != nil {
|
||||
t.Errorf("WriteGSO = %v, want nil", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkSegmentUDPv4 is the USO counterpart to BenchmarkSegmentTCPv4. The
|
||||
// yield is a no-op so the measurement is segmentation plus checksum work only.
|
||||
func BenchmarkSegmentUDPv4(b *testing.B) {
|
||||
sizes := []struct {
|
||||
name string
|
||||
payLen int
|
||||
gsoSize int
|
||||
}{
|
||||
{"64KiB_GSO1400", 64000, 1400},
|
||||
{"16KiB_GSO1400", 16384, 1400},
|
||||
{"4KiB_GSO1400", 4096, 1400},
|
||||
}
|
||||
for _, sz := range sizes {
|
||||
b.Run(sz.name, func(b *testing.B) {
|
||||
const ipLen = 20
|
||||
const udpLen = 8
|
||||
pkt := make([]byte, ipLen+udpLen+sz.payLen)
|
||||
pkt[0] = 0x45
|
||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+udpLen+sz.payLen))
|
||||
binary.BigEndian.PutUint16(pkt[4:6], 0x4242)
|
||||
pkt[8] = 64
|
||||
pkt[9] = unix.IPPROTO_UDP
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||
binary.BigEndian.PutUint16(pkt[20:22], 12345)
|
||||
binary.BigEndian.PutUint16(pkt[22:24], 53)
|
||||
binary.BigEndian.PutUint16(pkt[24:26], uint16(udpLen+sz.payLen))
|
||||
for i := 0; i < sz.payLen; i++ {
|
||||
pkt[ipLen+udpLen+i] = byte(i)
|
||||
}
|
||||
|
||||
master := append([]byte(nil), pkt...)
|
||||
work := make([]byte, len(pkt))
|
||||
p := Packet{Bytes: work, GSO: GSOInfo{
|
||||
Size: uint16(sz.gsoSize),
|
||||
HdrLen: ipLen + udpLen,
|
||||
CsumStart: ipLen,
|
||||
Proto: GSOProtoUDP,
|
||||
}}
|
||||
|
||||
b.SetBytes(int64(len(pkt)))
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
copy(work, master)
|
||||
if err := SegmentSuperpacket(p, func(seg []byte) error { return nil }); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkSegmentUDPv6 mirrors BenchmarkSegmentUDPv4 for IPv6, where the
|
||||
// pseudo-header address sum is 32 bytes rather than 8.
|
||||
func BenchmarkSegmentUDPv6(b *testing.B) {
|
||||
sizes := []struct {
|
||||
name string
|
||||
payLen int
|
||||
gsoSize int
|
||||
}{
|
||||
{"64KiB_GSO1400", 64000, 1400},
|
||||
{"16KiB_GSO1400", 16384, 1400},
|
||||
{"4KiB_GSO1400", 4096, 1400},
|
||||
}
|
||||
for _, sz := range sizes {
|
||||
b.Run(sz.name, func(b *testing.B) {
|
||||
const ipLen = 40
|
||||
const udpLen = 8
|
||||
pkt := make([]byte, ipLen+udpLen+sz.payLen)
|
||||
pkt[0] = 0x60
|
||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(udpLen+sz.payLen))
|
||||
pkt[6] = unix.IPPROTO_UDP
|
||||
pkt[7] = 64
|
||||
pkt[8], pkt[9], pkt[23] = 0xfe, 0x80, 1
|
||||
pkt[24], pkt[25], pkt[39] = 0xfe, 0x80, 2
|
||||
binary.BigEndian.PutUint16(pkt[40:42], 12345)
|
||||
binary.BigEndian.PutUint16(pkt[42:44], 53)
|
||||
binary.BigEndian.PutUint16(pkt[44:46], uint16(udpLen+sz.payLen))
|
||||
for i := 0; i < sz.payLen; i++ {
|
||||
pkt[ipLen+udpLen+i] = byte(i)
|
||||
}
|
||||
|
||||
master := append([]byte(nil), pkt...)
|
||||
work := make([]byte, len(pkt))
|
||||
p := Packet{Bytes: work, GSO: GSOInfo{
|
||||
Size: uint16(sz.gsoSize),
|
||||
HdrLen: ipLen + udpLen,
|
||||
CsumStart: ipLen,
|
||||
Proto: GSOProtoUDP,
|
||||
}}
|
||||
|
||||
b.SetBytes(int64(len(pkt)))
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
copy(work, master)
|
||||
if err := SegmentSuperpacket(p, func(seg []byte) error { return nil }); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,11 +3,7 @@
|
||||
|
||||
package virtio
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
import "encoding/binary"
|
||||
|
||||
// Size is the on-wire length of struct virtio_net_hdr the kernel
|
||||
// prepends/expects on a TUN opened with IFF_VNET_HDR (TUNSETVNETHDRSZ
|
||||
@@ -17,59 +13,31 @@ const Size = 10
|
||||
// Hdr is the Go view of the legacy virtio_net_hdr.
|
||||
type Hdr struct {
|
||||
Flags uint8
|
||||
gsoType uint8 //private to avoid mistakes wrt the 0x80 VIRTIO_NET_HDR_GSO_ECN flag, ORed with the other "GSO types"
|
||||
GSOType uint8
|
||||
HdrLen uint16
|
||||
GSOSize uint16
|
||||
CsumStart uint16
|
||||
CsumOffset uint16
|
||||
}
|
||||
|
||||
func NewHeader(flags, gsoType uint8, hdrLen, gsoSize, csumStart, csumOffset uint16) Hdr {
|
||||
return Hdr{
|
||||
Flags: flags,
|
||||
gsoType: gsoType,
|
||||
HdrLen: hdrLen,
|
||||
GSOSize: gsoSize,
|
||||
CsumStart: csumStart,
|
||||
CsumOffset: csumOffset,
|
||||
}
|
||||
}
|
||||
|
||||
// Decode reads a virtio_net_hdr in host byte order (TUN default; we never
|
||||
// call TUNSETVNETLE so the kernel matches our endianness).
|
||||
func (h *Hdr) Decode(b []byte) {
|
||||
h.Flags = b[0]
|
||||
h.gsoType = b[1]
|
||||
h.GSOType = b[1]
|
||||
h.HdrLen = binary.NativeEndian.Uint16(b[2:4])
|
||||
h.GSOSize = binary.NativeEndian.Uint16(b[4:6])
|
||||
h.CsumStart = binary.NativeEndian.Uint16(b[6:8])
|
||||
h.CsumOffset = binary.NativeEndian.Uint16(b[8:10])
|
||||
}
|
||||
|
||||
func EncodeHeader(b []byte, flags, gsoType uint8, hdrLen, gsoSize, csumStart, csumOffset uint16) {
|
||||
b[0] = flags
|
||||
b[1] = gsoType
|
||||
binary.NativeEndian.PutUint16(b[2:4], hdrLen)
|
||||
binary.NativeEndian.PutUint16(b[4:6], gsoSize)
|
||||
binary.NativeEndian.PutUint16(b[6:8], csumStart)
|
||||
binary.NativeEndian.PutUint16(b[8:10], csumOffset)
|
||||
}
|
||||
|
||||
// Encode is the inverse of Decode: writes the virtio_net_hdr fields into b
|
||||
// (must be at least Size bytes). Used to emit a TSO superpacket on egress.
|
||||
func (h *Hdr) Encode(b []byte) {
|
||||
EncodeHeader(b, h.Flags, h.gsoType, h.HdrLen, h.GSOSize, h.CsumStart, h.CsumOffset)
|
||||
}
|
||||
|
||||
// GSOType returns gsoType with the ECN-flag masked out
|
||||
func (h *Hdr) GSOType() uint8 {
|
||||
return h.gsoType &^ unix.VIRTIO_NET_HDR_GSO_ECN
|
||||
}
|
||||
|
||||
func (h *Hdr) HasECNFlag() bool {
|
||||
return h.gsoType&unix.VIRTIO_NET_HDR_GSO_ECN != 0
|
||||
}
|
||||
|
||||
func (h *Hdr) SetGSOType(x uint8) {
|
||||
h.gsoType = x
|
||||
b[0] = h.Flags
|
||||
b[1] = h.GSOType
|
||||
binary.NativeEndian.PutUint16(b[2:4], h.HdrLen)
|
||||
binary.NativeEndian.PutUint16(b[4:6], h.GSOSize)
|
||||
binary.NativeEndian.PutUint16(b[6:8], h.CsumStart)
|
||||
binary.NativeEndian.PutUint16(b[8:10], h.CsumOffset)
|
||||
}
|
||||
|
||||
+136
-143
@@ -27,8 +27,11 @@ const (
|
||||
tcpHeaderMaxLen = 60 // data-offset=15, max options
|
||||
)
|
||||
|
||||
// maxSegHdrLen bounds the L3+L4 header we snapshot before stamping each segment.
|
||||
// The largest header the segmenter supports is IPv4 (max IHL 60) plus TCP (max data-offset 60) = 120 bytes
|
||||
// maxSegHdrLen bounds the L3+L4 header we snapshot before stamping each
|
||||
// segment. The largest header the segmenter supports is IPv4 (max IHL 60)
|
||||
// plus TCP (max data-offset 60) = 120 bytes; the array is sized to that
|
||||
// worst case so the snapshot lives on the stack with no per-call heap
|
||||
// allocation.
|
||||
const maxSegHdrLen = ipv4HeaderMaxLen + tcpHeaderMaxLen // 120
|
||||
|
||||
// Byte offsets inside an IPv4 header.
|
||||
@@ -62,72 +65,63 @@ const (
|
||||
udpChecksumOff = 6
|
||||
)
|
||||
|
||||
var errPacketTooShort = errors.New("packet too short")
|
||||
|
||||
// tcpFinPshMask is cleared on every segment except the last of a TSO burst.
|
||||
const tcpFinPshMask = 0x09 // FIN(0x01) | PSH(0x08)
|
||||
|
||||
// tcpCwrFlag is cleared on every segment except the first.
|
||||
// Per RFC 3168 §6.1.2 the CWR bit signals a one-shot transition (the sender just halved its window)
|
||||
// and must appear on the first segment of a TSO burst only.
|
||||
// tcpCwrFlag is cleared on every segment except the first. Per RFC 3168
|
||||
// §6.1.2 the CWR bit signals a one-shot transition (the sender just halved
|
||||
// its window) and must appear on the first segment of a TSO burst only.
|
||||
const tcpCwrFlag = 0x80
|
||||
|
||||
// CheckValid rejects packets whose virtio_net_hdr/IP combination would
|
||||
// cause a downstream miscompute. The TUN should never emit RSC_INFO and
|
||||
// the GSO type must agree with the IP version nibble.
|
||||
func CheckValid(pkt []byte, hdr Hdr) error {
|
||||
// When RSC_INFO is set the csum_start/csum_offset fields are repurposed to
|
||||
// carry coalescing info rather than checksum offsets. A TUN writing via
|
||||
// IFF_VNET_HDR should never emit this, but if it did we would silently
|
||||
// miscompute the segment checksums — refuse the packet instead.
|
||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_RSC_INFO != 0 {
|
||||
return fmt.Errorf("virtio RSC_INFO flag not supported on TUN reads")
|
||||
}
|
||||
if len(pkt) < ipv4HeaderMinLen {
|
||||
return errPacketTooShort
|
||||
return fmt.Errorf("packet too short")
|
||||
}
|
||||
ipVersion := pkt[0] >> 4
|
||||
if ipVersion == 6 && len(pkt) < ipv6FixedLen {
|
||||
return errPacketTooShort
|
||||
}
|
||||
|
||||
gsoType := hdr.GSOType()
|
||||
if gsoType != unix.VIRTIO_NET_HDR_GSO_NONE && hdr.GSOSize == 0 {
|
||||
// A GSO type with no segment size would dodge IsSuperpacket() downstream and
|
||||
// travel as a plain jumbo datagram with an unfinished checksum.
|
||||
return fmt.Errorf("virtio GSO type %#x with zero gso_size", hdr.gsoType)
|
||||
}
|
||||
if hdr.HasECNFlag() && !(gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV4 || gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV6) {
|
||||
return fmt.Errorf("virtio GSO_ECN qualifier on non-TCP GSO type %#x", hdr.gsoType)
|
||||
}
|
||||
switch gsoType {
|
||||
switch hdr.GSOType {
|
||||
case unix.VIRTIO_NET_HDR_GSO_TCPV4:
|
||||
if ipVersion != 4 {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
||||
}
|
||||
case unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
||||
if ipVersion != 6 {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
||||
}
|
||||
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
|
||||
// USO carries either v4 or v6; the leading nibble disambiguates.
|
||||
if !(ipVersion == 4 || ipVersion == 6) {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
||||
}
|
||||
default:
|
||||
if !(ipVersion == 6 || ipVersion == 4) {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// CorrectHdrLen rewrites hdr.HdrLen based on the actual transport header length read out of pkt.
|
||||
// The kernel's hdr.HdrLen on the FORWARD path can be the length of the entire first packet, so we don't trust it.
|
||||
// CorrectHdrLen rewrites hdr.HdrLen based on the actual transport header
|
||||
// length read out of pkt. The kernel's hdr.HdrLen on the FORWARD path can
|
||||
// be the length of the entire first packet, so we don't trust it.
|
||||
func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
|
||||
// Thank you wireguard-go for documenting these edge-cases
|
||||
// Don't trust hdr.hdrLen from the kernel as it can be equal to the length
|
||||
// of the entire first packet when the kernel is handling it as part of a FORWARD path.
|
||||
// Instead, parse the transport header length and add it onto csumStart, which is synonymous for IP header length.
|
||||
// of the entire first packet when the kernel is handling it as part of a
|
||||
// FORWARD path. Instead, parse the transport header length and add it onto
|
||||
// csumStart, which is synonymous for IP header length.
|
||||
|
||||
if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
|
||||
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
|
||||
hdr.HdrLen = hdr.CsumStart + 8
|
||||
} else {
|
||||
if len(pkt) <= int(hdr.CsumStart+tcpDataOffOff) {
|
||||
@@ -135,7 +129,8 @@ func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
|
||||
}
|
||||
|
||||
tcpHLen := uint16(pkt[hdr.CsumStart+tcpDataOffOff] >> 4 * 4)
|
||||
if tcpHLen < tcpHeaderMinLen || tcpHLen > tcpHeaderMaxLen {
|
||||
if tcpHLen < 20 || tcpHLen > 60 {
|
||||
// A TCP header must be between 20 and 60 bytes in length.
|
||||
return fmt.Errorf("tcp header len is invalid: %d", tcpHLen)
|
||||
}
|
||||
hdr.HdrLen = hdr.CsumStart + tcpHLen
|
||||
@@ -155,62 +150,19 @@ func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// segCount returns how many segments a payload of payLen bytes splits into at gsoSize,
|
||||
// with a floor of one so a header-only superpacket still yields a single segment.
|
||||
func segCount(payLen, gsoSize int) int {
|
||||
n := (payLen + gsoSize - 1) / gsoSize
|
||||
if n == 0 {
|
||||
return 1
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
// basePseudoSum folds the part of the L4 pseudo-header sum that is identical
|
||||
// for every segment: the source and destination addresses plus the protocol
|
||||
// number. The per-segment L4 length is added by the caller inside the loop.
|
||||
func basePseudoSum(pkt []byte, isV4 bool, proto uint32) uint32 {
|
||||
if isV4 {
|
||||
return uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0)) + proto
|
||||
}
|
||||
return uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0)) + proto
|
||||
}
|
||||
|
||||
// baseIPv4HdrSum folds the IPv4 header checksum over the fields that stay constant across segments.
|
||||
// csumStart is the L3 header length, which bounds a valid IHL.
|
||||
func baseIPv4HdrSum(pkt []byte, csumStart int) (uint32, error) {
|
||||
ihl := int(pkt[0]&0x0f) * 4
|
||||
if ihl < ipv4HeaderMinLen || ihl > csumStart {
|
||||
return 0, fmt.Errorf("bad IPv4 IHL: %d", ihl)
|
||||
}
|
||||
// total_len, the ID, and the checksum field itself are excluded: all three are rewritten per segment.
|
||||
sum := uint32(checksum.Checksum(pkt[:ihl], 0))
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4TotalLenOff : ipv4TotalLenOff+2]))
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4ChecksumOff : ipv4ChecksumOff+2]))
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2]))
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
return sum, nil
|
||||
}
|
||||
|
||||
// baseTCPHdrSum folds the TCP header checksum over everything the segment loop does not rewrite
|
||||
func baseTCPHdrSum(pkt []byte, csumStart, headerLen int) uint32 {
|
||||
seq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
|
||||
flags := uint16(pkt[csumStart+tcpFlagsOff])
|
||||
|
||||
sum := uint32(checksum.Checksum(pkt[csumStart:headerLen], 0))
|
||||
sum += uint32(^uint16(seq >> 16))
|
||||
sum += uint32(^uint16(seq))
|
||||
sum += uint32(^flags)
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[csumStart+tcpChecksumOff : csumStart+tcpChecksumOff+2]))
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
return sum
|
||||
}
|
||||
|
||||
// SegmentTCP walks a TSO superpacket pkt, yielding each segment as a slice into pkt.
|
||||
// Per-segment plaintext is laid out by stamping a copy of the original L3+L4 header into pkt at offset i*gsoSize,
|
||||
// where it sits immediately before that segment's payload chunk in the original buffer.
|
||||
// pkt is consumed by this call and must not be inspected by the caller after the final yield.
|
||||
// SegmentTCP walks a TSO superpacket pkt, yielding each segment as a
|
||||
// slice into pkt itself. Per-segment plaintext is laid out by stamping a
|
||||
// copy of the original L3+L4 header into pkt at offset i*gsoSize, where it
|
||||
// sits immediately before that segment's payload chunk in the original
|
||||
// buffer. The stamp is destructive but harmless: iter i's header write lands
|
||||
// on pkt[i*G : i*G+hdrLen], which is the tail of seg_{i-1}'s payload (already
|
||||
// consumed) and ends exactly where seg_i's payload begins, so it never clobbers
|
||||
// live payload — this holds even when gsoSize < hdrLen. The header bytes are
|
||||
// sourced from a pristine snapshot taken before the loop (savedHdr), NOT from
|
||||
// pkt[:hdrLen], because when gsoSize < hdrLen the stamps would otherwise
|
||||
// overwrite the leading header in place and every stamp after the first would
|
||||
// copy corrupted bytes. pkt is consumed by this call and must not be inspected
|
||||
// by the caller after the final yield.
|
||||
func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
|
||||
if gsoSizeU == 0 {
|
||||
return fmt.Errorf("gso_size is zero")
|
||||
@@ -229,28 +181,49 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
||||
tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4
|
||||
payLen := len(pkt) - headerLen
|
||||
gsoSize := int(gsoSizeU)
|
||||
numSeg := segCount(payLen, gsoSize)
|
||||
numSeg := (payLen + gsoSize - 1) / gsoSize
|
||||
if numSeg == 0 {
|
||||
numSeg = 1
|
||||
}
|
||||
|
||||
origSeq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
|
||||
origFlags := pkt[csumStart+tcpFlagsOff]
|
||||
|
||||
baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_TCP)
|
||||
baseTcpHdrSum := baseTCPHdrSum(pkt, csumStart, headerLen)
|
||||
var tmp [tcpHeaderMaxLen]byte
|
||||
copy(tmp[:tcpHdrLen], pkt[csumStart:headerLen])
|
||||
tmp[tcpSeqOff], tmp[tcpSeqOff+1], tmp[tcpSeqOff+2], tmp[tcpSeqOff+3] = 0, 0, 0, 0
|
||||
tmp[tcpFlagsOff] = 0
|
||||
tmp[tcpChecksumOff], tmp[tcpChecksumOff+1] = 0, 0
|
||||
baseTcpHdrSum := uint32(checksum.Checksum(tmp[:tcpHdrLen], 0))
|
||||
|
||||
var baseProtoSum uint32
|
||||
if isV4 {
|
||||
baseProtoSum = uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0))
|
||||
} else {
|
||||
baseProtoSum = uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0))
|
||||
}
|
||||
baseProtoSum += uint32(unix.IPPROTO_TCP)
|
||||
|
||||
var origIPID uint16
|
||||
var baseIPHdrSum uint32
|
||||
if isV4 {
|
||||
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
|
||||
var err error
|
||||
// TSO bumps the ID per segment, so it stays out of the base sum.
|
||||
baseIPHdrSum, err = baseIPv4HdrSum(pkt, csumStart)
|
||||
if err != nil {
|
||||
return err
|
||||
ihl := int(pkt[0]&0x0f) * 4
|
||||
if ihl < ipv4HeaderMinLen || ihl > csumStart {
|
||||
return fmt.Errorf("bad IPv4 IHL: %d", ihl)
|
||||
}
|
||||
var ipTmp [ipv4HeaderMaxLen]byte
|
||||
copy(ipTmp[:ihl], pkt[:ihl])
|
||||
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
|
||||
ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0
|
||||
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
|
||||
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
|
||||
}
|
||||
|
||||
// Snapshot the pristine L3+L4 header once. '
|
||||
// Every segment's header is stamped from this copy, so overlapping stamps (gsoSize < headerLen) can never corrupt the source.
|
||||
// Snapshot the pristine L3+L4 header once. Every segment's header is
|
||||
// stamped from this copy, so overlapping stamps (gsoSize < headerLen)
|
||||
// can never corrupt the source. The variable fields (seq/flags/cksum/
|
||||
// totalLen/id) captured here are stale but are overwritten per segment.
|
||||
var savedHdr [maxSegHdrLen]byte
|
||||
copy(savedHdr[:headerLen], pkt[:headerLen])
|
||||
|
||||
@@ -264,10 +237,12 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
||||
segLen := headerLen + segPayLen
|
||||
headerOff := i * gsoSize
|
||||
|
||||
// Stamp the header into place immediately before this segment's payload, sourced from the snapshot.
|
||||
// The per-segment patches below overwrite the variable fields. (seq/flags/cksum/totalLen/id)
|
||||
// Stamp the header into place immediately before this segment's
|
||||
// payload, sourced from the pristine snapshot. Iter 0's header is
|
||||
// already at pkt[:headerLen] (identical to savedHdr), so only i ≥ 1
|
||||
// needs the stamp. The per-segment patches below overwrite the
|
||||
// variable fields.
|
||||
if i > 0 {
|
||||
// Iter 0's header is already at pkt[:headerLen] (identical to savedHdr), so only i >= 1 needs the stamp
|
||||
copy(pkt[headerOff:headerOff+headerLen], savedHdr[:headerLen])
|
||||
}
|
||||
seg := pkt[headerOff : headerOff+segLen]
|
||||
@@ -296,9 +271,10 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
||||
seg[csumStart+tcpFlagsOff] = segFlags
|
||||
|
||||
tcpLen := tcpHdrLen + segPayLen
|
||||
// Payload bytes still live at their original offset in pkt.
|
||||
// The header slide above only writes into pkt[i*GSOSize : i*GSOSize+header], which is the tail of seg_{i-1}'s payload (already consumed)
|
||||
// and never overlaps seg_i's own payload at pkt[header+i*GSOSize : header+(i+1)*GSOSize].
|
||||
// Payload bytes still live at their original offset in pkt. The
|
||||
// header slide above only writes into pkt[i*G : i*G+H], which is
|
||||
// the tail of seg_{i-1}'s payload (already consumed) and never
|
||||
// overlaps seg_i's own payload at pkt[H+i*G : H+(i+1)*G].
|
||||
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
|
||||
wide := uint64(baseTcpHdrSum) + uint64(paySum) + uint64(baseProtoSum)
|
||||
wide += uint64(segSeq) + uint64(segFlags) + uint64(tcpLen)
|
||||
@@ -314,10 +290,17 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
||||
return nil
|
||||
}
|
||||
|
||||
// SegmentUDP walks a USO superpacket, stamping a per-segment-patched copy of the original L3+L4 header
|
||||
// into pkt at offset i*GSOSize and yielding pkt[i*GSOSize:i*GSOSize+segLen] to the caller.
|
||||
// Per-segment patches are total_len + IPv4 csum (or IPv6 payload_len) plus the UDP length and checksum.
|
||||
// pkt is consumed destructively.
|
||||
// SegmentUDP walks a USO superpacket, stamping a per-segment-patched copy of
|
||||
// the original L3+L4 header into pkt at offset i*gsoSize and yielding
|
||||
// pkt[i*G:i*G+segLen] to the caller. Per-segment patches are total_len +
|
||||
// IPv4 csum (or IPv6 payload_len) plus the UDP length and checksum. pkt is
|
||||
// consumed destructively; see SegmentTCP for the layout reasoning, including
|
||||
// why the header is stamped from a pristine snapshot rather than pkt[:hdrLen]
|
||||
// (correctness when gsoSize < hdrLen).
|
||||
//
|
||||
// UDP-GSO leaves the IPv4 ID identical across segments (the kernel does not
|
||||
// bump it), which is why the IP-level per-segment work is limited to
|
||||
// total_len + IPv4 header checksum (v4) or payload_len (v6).
|
||||
func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
|
||||
if gsoSizeU == 0 {
|
||||
return fmt.Errorf("gso_size is zero")
|
||||
@@ -338,24 +321,41 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
||||
|
||||
payLen := len(pkt) - headerLen
|
||||
gsoSize := int(gsoSizeU)
|
||||
numSeg := segCount(payLen, gsoSize)
|
||||
|
||||
baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_UDP)
|
||||
|
||||
var origIPID uint16
|
||||
var baseIPHdrSum uint32
|
||||
if isV4 {
|
||||
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
|
||||
var err error
|
||||
// Software UDP GSO bumps the ID per segment just like TSO
|
||||
// (inet_gso_segment's fixed-ID case is TCP-only), so it stays out of the base sum.
|
||||
baseIPHdrSum, err = baseIPv4HdrSum(pkt, csumStart)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
numSeg := (payLen + gsoSize - 1) / gsoSize
|
||||
if numSeg == 0 {
|
||||
numSeg = 1
|
||||
}
|
||||
|
||||
// Snapshot the pristine L3+L4 header once and stamp every segment from it
|
||||
var udpTmp [udpHeaderLen]byte
|
||||
copy(udpTmp[:], pkt[csumStart:headerLen])
|
||||
udpTmp[udpLengthOff], udpTmp[udpLengthOff+1] = 0, 0
|
||||
udpTmp[udpChecksumOff], udpTmp[udpChecksumOff+1] = 0, 0
|
||||
baseUDPHdrSum := uint32(checksum.Checksum(udpTmp[:], 0))
|
||||
|
||||
var baseProtoSum uint32
|
||||
if isV4 {
|
||||
baseProtoSum = uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0))
|
||||
} else {
|
||||
baseProtoSum = uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0))
|
||||
}
|
||||
baseProtoSum += uint32(unix.IPPROTO_UDP)
|
||||
|
||||
var baseIPHdrSum uint32
|
||||
if isV4 {
|
||||
ihl := int(pkt[0]&0x0f) * 4
|
||||
if ihl < ipv4HeaderMinLen || ihl > csumStart {
|
||||
return fmt.Errorf("bad IPv4 IHL: %d", ihl)
|
||||
}
|
||||
var ipTmp [ipv4HeaderMaxLen]byte
|
||||
copy(ipTmp[:ihl], pkt[:ihl])
|
||||
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
|
||||
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
|
||||
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
|
||||
}
|
||||
|
||||
// Snapshot the pristine L3+L4 header once and stamp every segment from
|
||||
// it; see SegmentTCP for why sourcing from pkt[:headerLen] corrupts
|
||||
// segments when gsoSize < headerLen.
|
||||
var savedHdr [maxSegHdrLen]byte
|
||||
copy(savedHdr[:headerLen], pkt[:headerLen])
|
||||
|
||||
@@ -378,10 +378,8 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
||||
udpLen := udpHeaderLen + segPayLen
|
||||
|
||||
if isV4 {
|
||||
segID := origIPID + uint16(i)
|
||||
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
|
||||
binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID)
|
||||
ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID)
|
||||
ipSum := baseIPHdrSum + uint32(totalLen)
|
||||
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
|
||||
} else {
|
||||
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
|
||||
@@ -389,13 +387,12 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
||||
|
||||
binary.BigEndian.PutUint16(seg[csumStart+udpLengthOff:csumStart+udpLengthOff+2], uint16(udpLen))
|
||||
|
||||
// Sum the UDP header (length just written, checksum zeroed) together with
|
||||
// this segment's payload in one pass, seeded with the pseudo-header sum.
|
||||
seg[csumStart+udpChecksumOff], seg[csumStart+udpChecksumOff+1] = 0, 0
|
||||
pseudo := baseProtoSum + uint32(udpLen)
|
||||
pseudo = (pseudo & 0xffff) + (pseudo >> 16)
|
||||
pseudo = (pseudo & 0xffff) + (pseudo >> 16)
|
||||
csum := ^checksum.Checksum(seg[csumStart:], uint16(pseudo))
|
||||
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
|
||||
wide := uint64(baseUDPHdrSum) + uint64(paySum) + uint64(baseProtoSum)
|
||||
wide += uint64(udpLen) + uint64(udpLen)
|
||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
||||
csum := foldComplement(uint32(wide))
|
||||
if csum == 0 {
|
||||
csum = 0xffff
|
||||
}
|
||||
@@ -409,9 +406,10 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
|
||||
return nil
|
||||
}
|
||||
|
||||
// FinishChecksum computes the L4 checksum for a non-GSO packet that the kernel handed us with NEEDS_CSUM set.
|
||||
// CsumStart / CsumOffset point at the 16-bit checksum field.
|
||||
// We zero it, fold a full sum from the partial one that the kernel provided, and store the result.
|
||||
// FinishChecksum computes the L4 checksum for a non-GSO packet that the kernel
|
||||
// handed us with NEEDS_CSUM set. csum_start / csum_offset point at the 16-bit
|
||||
// checksum field; we zero it, fold a full sum (the field was pre-loaded with
|
||||
// the pseudo-header partial sum by the kernel), and store the result.
|
||||
func FinishChecksum(seg []byte, hdr Hdr) error {
|
||||
cs := int(hdr.CsumStart)
|
||||
co := int(hdr.CsumOffset)
|
||||
@@ -423,12 +421,7 @@ func FinishChecksum(seg []byte, hdr Hdr) error {
|
||||
partial := binary.BigEndian.Uint16(seg[cs+co : cs+co+2])
|
||||
seg[cs+co] = 0
|
||||
seg[cs+co+1] = 0
|
||||
csum := ^checksum.Checksum(seg[cs:], partial)
|
||||
// RFC 768: UDP transmits a computed zero as all ones, since all-zero is the reserved "no checksum" value.
|
||||
if co == udpChecksumOff && csum == 0 {
|
||||
csum = 0xffff
|
||||
}
|
||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], csum)
|
||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], ^checksum.Checksum(seg[cs:], partial))
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -226,14 +226,13 @@ func TestCorrectHdrLenChecksumBound(t *testing.T) {
|
||||
// = 26) accepts. This case FAILS against the CsumStart+CsumStart regression.
|
||||
t.Run("valid-small-uso-accepted", func(t *testing.T) {
|
||||
pkt, _, csumStart := buildUDPv4Super(12) // total len 40
|
||||
hdr := NewHeader(
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
|
||||
0, /*hdrLen*/
|
||||
6, /*gsoSize: two 6-byte segments*/
|
||||
csumStart, /*csumStart*/
|
||||
6, /*csumOffset*/
|
||||
)
|
||||
hdr := Hdr{
|
||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
|
||||
GSOSize: 6, // two 6-byte segments
|
||||
CsumStart: csumStart,
|
||||
CsumOffset: 6,
|
||||
}
|
||||
if err := CorrectHdrLen(pkt, &hdr); err != nil {
|
||||
t.Fatalf("CorrectHdrLen rejected a valid 40-byte USO superpacket: %v", err)
|
||||
}
|
||||
@@ -248,14 +247,13 @@ func TestCorrectHdrLenChecksumBound(t *testing.T) {
|
||||
t.Run("too-short-rejected", func(t *testing.T) {
|
||||
pkt := make([]byte, 25)
|
||||
pkt[0] = 0x45 // IPv4, IHL 5
|
||||
hdr := NewHeader(
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
|
||||
0, /*hdrLen*/
|
||||
6, /*gsoSize*/
|
||||
20, /*csumStart*/
|
||||
6, /*csumOffset*/
|
||||
)
|
||||
hdr := Hdr{
|
||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
|
||||
GSOSize: 6,
|
||||
CsumStart: 20,
|
||||
CsumOffset: 6,
|
||||
}
|
||||
if err := CorrectHdrLen(pkt, &hdr); err == nil {
|
||||
t.Fatalf("CorrectHdrLen accepted a too-short (25-byte) packet")
|
||||
}
|
||||
@@ -305,10 +303,9 @@ func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
|
||||
if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 53 {
|
||||
t.Errorf("seg %d: dport=%d want 53", i, dport)
|
||||
}
|
||||
// Software UDP GSO bumps the IPv4 ID per segment just like TSO
|
||||
// (inet_gso_segment's fixed-ID case is TCP-only).
|
||||
if id := binary.BigEndian.Uint16(seg[4:6]); id != 0x4242+uint16(i) {
|
||||
t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242+uint16(i))
|
||||
// UDP-GSO keeps the same IPv4 ID across every segment.
|
||||
if id := binary.BigEndian.Uint16(seg[4:6]); id != 0x4242 {
|
||||
t.Errorf("seg %d: ip id=%#x want 0x4242", i, id)
|
||||
}
|
||||
|
||||
segPayLen := len(seg) - int(hdrLen)
|
||||
@@ -336,267 +333,3 @@ func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// buildUDPv4Single constructs a single IPv4/UDP datagram with the checksum field pre-loaded with the folded
|
||||
// pseudo-header sum, which is how a NEEDS_CSUM (CHECKSUM_PARTIAL) packet arrives from the tun.
|
||||
func buildUDPv4Single(payload []byte) (pkt []byte, hdr Hdr) {
|
||||
const ipLen, udpLen = 20, 8
|
||||
pkt = make([]byte, ipLen+udpLen+len(payload))
|
||||
|
||||
pkt[0] = 0x45
|
||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(len(pkt)))
|
||||
pkt[8] = 64
|
||||
pkt[9] = unix.IPPROTO_UDP
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||
|
||||
binary.BigEndian.PutUint16(pkt[ipLen:ipLen+2], 12345)
|
||||
binary.BigEndian.PutUint16(pkt[ipLen+2:ipLen+4], 53)
|
||||
binary.BigEndian.PutUint16(pkt[ipLen+4:ipLen+6], uint16(udpLen+len(payload)))
|
||||
copy(pkt[ipLen+udpLen:], payload)
|
||||
|
||||
pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, udpLen+len(payload))
|
||||
binary.BigEndian.PutUint16(pkt[ipLen+udpChecksumOff:ipLen+udpChecksumOff+2], pseudo)
|
||||
|
||||
return pkt, NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, ipLen, udpChecksumOff)
|
||||
}
|
||||
|
||||
// TestFinishChecksumUDPZeroStoresAllOnes pins RFC 768: a UDP checksum that computes to zero goes on the wire as
|
||||
// 0xffff, because all-zero is the reserved "no checksum" encoding that IPv6 rejects outright.
|
||||
func TestFinishChecksumUDPZeroStoresAllOnes(t *testing.T) {
|
||||
var payload []byte
|
||||
for i := 0; i < 0x10000; i++ {
|
||||
p := []byte{byte(i >> 8), byte(i)}
|
||||
pkt, hdr := buildUDPv4Single(p)
|
||||
cs, co := int(hdr.CsumStart), int(hdr.CsumOffset)
|
||||
partial := binary.BigEndian.Uint16(pkt[cs+co : cs+co+2])
|
||||
pkt[cs+co], pkt[cs+co+1] = 0, 0
|
||||
if ^checksum.Checksum(pkt[cs:], partial) == 0 {
|
||||
payload = p
|
||||
break
|
||||
}
|
||||
}
|
||||
if payload == nil {
|
||||
t.Fatal("no 2-byte payload produced a zero checksum")
|
||||
}
|
||||
|
||||
pkt, hdr := buildUDPv4Single(payload)
|
||||
if err := FinishChecksum(pkt, hdr); err != nil {
|
||||
t.Fatalf("FinishChecksum: %v", err)
|
||||
}
|
||||
off := int(hdr.CsumStart) + int(hdr.CsumOffset)
|
||||
if got := binary.BigEndian.Uint16(pkt[off : off+2]); got != 0xffff {
|
||||
t.Fatalf("stored %#04x, want 0xffff for a UDP checksum that folds to zero", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFinishChecksumTCPZeroPreserved is the other half of the protocol split: zero is a legal TCP checksum and must
|
||||
// be stored as-is, so the UDP rewrite must not fire on a TCP csum_offset.
|
||||
func TestFinishChecksumTCPZeroPreserved(t *testing.T) {
|
||||
const cs, co = 20, tcpChecksumOff
|
||||
|
||||
// Seed the partial so the completed checksum lands on zero, the case the UDP path rewrites.
|
||||
seg := make([]byte, cs+co+2)
|
||||
for i := range seg[cs:] {
|
||||
seg[cs+i] = byte(i * 7)
|
||||
}
|
||||
var partial uint16
|
||||
for i := 0; i <= 0xffff; i++ {
|
||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], uint16(i))
|
||||
probe := append([]byte(nil), seg...)
|
||||
probe[cs+co], probe[cs+co+1] = 0, 0
|
||||
if ^checksum.Checksum(probe[cs:], uint16(i)) == 0 {
|
||||
partial = uint16(i)
|
||||
break
|
||||
}
|
||||
}
|
||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], partial)
|
||||
|
||||
hdr := NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, cs, co)
|
||||
if err := FinishChecksum(seg, hdr); err != nil {
|
||||
t.Fatalf("FinishChecksum: %v", err)
|
||||
}
|
||||
if got := binary.BigEndian.Uint16(seg[cs+co : cs+co+2]); got != 0 {
|
||||
t.Fatalf("stored %#04x, want 0x0000 (zero is a legal TCP checksum)", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFinishChecksumUDPValidates confirms the ordinary path still produces a checksum a receiver accepts.
|
||||
func TestFinishChecksumUDPValidates(t *testing.T) {
|
||||
payload := []byte("the definitive tun offloads branch")
|
||||
pkt, hdr := buildUDPv4Single(payload)
|
||||
if err := FinishChecksum(pkt, hdr); err != nil {
|
||||
t.Fatalf("FinishChecksum: %v", err)
|
||||
}
|
||||
pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, 8+len(payload))
|
||||
if !verifyChecksum(pkt[hdr.CsumStart:], pseudo) {
|
||||
t.Fatal("completed UDP checksum does not validate")
|
||||
}
|
||||
}
|
||||
|
||||
// TestCheckValidMasksGSOECN: GSO_ECN is a qualifier bit the kernel ORs
|
||||
// into gso_type for TSO superpackets with CWR set. CheckValid must
|
||||
// validate an ECN-qualified type as its base type — previously TCPV4|ECN
|
||||
// fell into the default case and skipped the IP-version agreement check.
|
||||
// The qualifier is TCP-only, so it must be rejected on UDP_L4.
|
||||
func TestCheckValidMasksGSOECN(t *testing.T) {
|
||||
v4pkt, _, _ := buildTCPv4Super(100)
|
||||
v6pkt := make([]byte, len(v4pkt))
|
||||
copy(v6pkt, v4pkt)
|
||||
v6pkt[0] = 0x60 // claim IPv6
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
pkt []byte
|
||||
gsoType uint8
|
||||
wantErr bool
|
||||
}{
|
||||
{"tcpv4-ecn-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, false},
|
||||
{"tcpv4-ecn-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
|
||||
{"tcpv6-ecn-v4-mismatch", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV6 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
|
||||
{"udp-l4-ecn-rejected", v4pkt, unix.VIRTIO_NET_HDR_GSO_UDP_L4 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
|
||||
{"tcpv4-plain-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, false},
|
||||
{"tcpv4-plain-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, true},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
err := CheckValid(tc.pkt, NewHeader(0, tc.gsoType, 0, 100, 0, 0))
|
||||
if tc.wantErr && err == nil {
|
||||
t.Errorf("CheckValid(gsoType=%#x) = nil, want error", tc.gsoType)
|
||||
}
|
||||
if !tc.wantErr && err != nil {
|
||||
t.Errorf("CheckValid(gsoType=%#x) = %v, want nil", tc.gsoType, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestCheckValidRejectsZeroGSOSize: a GSO-typed header with gso_size=0 must be
|
||||
// rejected. It would produce a Packet whose GSOInfo.IsSuperpacket() is false,
|
||||
// dodging both segmentation and FinishChecksum on its way downstream.
|
||||
func TestCheckValidRejectsZeroGSOSize(t *testing.T) {
|
||||
v4pkt, _, _ := buildTCPv4Super(100)
|
||||
if err := CheckValid(v4pkt, NewHeader(0, unix.VIRTIO_NET_HDR_GSO_TCPV4, 0, 0, 0, 0)); err == nil {
|
||||
t.Fatal("CheckValid accepted a GSO-typed header with gso_size=0")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFoldComplementMatchesReference checks the segmenter's fold-and-invert
|
||||
// against an independent RFC 1071 reference fold, hitting the carry edge
|
||||
// cases (values whose first fold produces another carry).
|
||||
func TestFoldComplementMatchesReference(t *testing.T) {
|
||||
refFold := func(s uint64) uint16 {
|
||||
for s>>16 != 0 {
|
||||
s = s&0xffff + s>>16
|
||||
}
|
||||
return uint16(s)
|
||||
}
|
||||
cases := []uint32{
|
||||
0, 1, 0xffff,
|
||||
0x10000, // single carry
|
||||
0x1fffe, // 0xffff + 0xffff: carry produces another 0xffff
|
||||
0xffff0000, // high half only
|
||||
0xfffeffff, // first fold yields another carry
|
||||
0xffffffff, // worst case
|
||||
}
|
||||
for _, c := range cases {
|
||||
if got, want := foldComplement(c), ^refFold(uint64(c)); got != want {
|
||||
t.Errorf("foldComplement(%#x) = %#x, want %#x", c, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// referenceBaseIPv4HdrSum and referenceBaseTCPHdrSum are the straightforward
|
||||
// implementations that baseIPv4HdrSum/baseTCPHdrSum replaced: copy the header
|
||||
// into scratch, zero the fields the segment loop rewrites, sum. The production
|
||||
// versions instead sum in place and subtract those fields via one's-complement
|
||||
// arithmetic, which is faster but far less obvious — particularly for the TCP
|
||||
// flags byte, which is only half of a 16-bit word. These references exist so
|
||||
// that trade is checked rather than asserted.
|
||||
func referenceBaseIPv4HdrSum(pkt []byte, ihl int) uint32 {
|
||||
var ipTmp [ipv4HeaderMaxLen]byte
|
||||
copy(ipTmp[:ihl], pkt[:ihl])
|
||||
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
|
||||
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
|
||||
ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0
|
||||
return uint32(checksum.Checksum(ipTmp[:ihl], 0))
|
||||
}
|
||||
|
||||
func referenceBaseTCPHdrSum(pkt []byte, csumStart, headerLen int) uint32 {
|
||||
tcpLen := headerLen - csumStart
|
||||
var tmp [tcpHeaderMaxLen]byte
|
||||
copy(tmp[:tcpLen], pkt[csumStart:headerLen])
|
||||
tmp[tcpSeqOff], tmp[tcpSeqOff+1], tmp[tcpSeqOff+2], tmp[tcpSeqOff+3] = 0, 0, 0, 0
|
||||
tmp[tcpFlagsOff] = 0
|
||||
tmp[tcpChecksumOff], tmp[tcpChecksumOff+1] = 0, 0
|
||||
return uint32(checksum.Checksum(tmp[:tcpLen], 0))
|
||||
}
|
||||
|
||||
// randSeed is a tiny deterministic PRNG so this test needs no imports beyond
|
||||
// what the file already has and reproduces identically on every run.
|
||||
func randByte(state *uint32) byte {
|
||||
*state = *state*1664525 + 1013904223
|
||||
return byte(*state >> 24)
|
||||
}
|
||||
|
||||
func TestBaseSumsMatchZeroingReference(t *testing.T) {
|
||||
state := uint32(12345)
|
||||
|
||||
t.Run("ipv4", func(t *testing.T) {
|
||||
for ihl := ipv4HeaderMinLen; ihl <= ipv4HeaderMaxLen; ihl += 4 {
|
||||
for iter := 0; iter < 5000; iter++ {
|
||||
pkt := make([]byte, ihl)
|
||||
for i := range pkt {
|
||||
pkt[i] = randByte(&state)
|
||||
}
|
||||
pkt[0] = byte(0x40 | (ihl / 4))
|
||||
|
||||
want := referenceBaseIPv4HdrSum(pkt, ihl)
|
||||
got, err := baseIPv4HdrSum(pkt, ihl)
|
||||
if err != nil {
|
||||
t.Fatalf("ihl=%d: %v", ihl, err)
|
||||
}
|
||||
// Compare the value that reaches the wire: the raw partial
|
||||
// sums may legally differ by one's-complement -0 vs +0.
|
||||
for _, tl := range []uint32{20, 1500, 65535} {
|
||||
for _, id := range []uint32{0, 0x4242, 0xffff} {
|
||||
if a, b := foldComplement(want+tl+id), foldComplement(got+tl+id); a != b {
|
||||
t.Fatalf("ihl=%d tl=%d id=%d: %#04x != %#04x", ihl, tl, id, a, b)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("tcp", func(t *testing.T) {
|
||||
const csumStart = 20
|
||||
for dataOff := 5; dataOff <= 15; dataOff++ {
|
||||
tcpLen := dataOff * 4
|
||||
headerLen := csumStart + tcpLen
|
||||
for iter := 0; iter < 5000; iter++ {
|
||||
pkt := make([]byte, headerLen+64)
|
||||
for i := range pkt {
|
||||
pkt[i] = randByte(&state)
|
||||
}
|
||||
pkt[0] = 0x45
|
||||
pkt[csumStart+tcpDataOffOff] = byte(dataOff << 4)
|
||||
|
||||
want := referenceBaseTCPHdrSum(pkt, csumStart, headerLen)
|
||||
got := baseTCPHdrSum(pkt, csumStart, headerLen)
|
||||
for _, seq := range []uint32{0, 1, 0x4242_4242, 0xffff_ffff} {
|
||||
for _, fl := range []uint32{0x00, 0x10, 0x18, 0x19, 0xff} {
|
||||
for _, l4 := range []uint32{20, 1460, 65535} {
|
||||
a := foldComplement(want + seq + fl + l4)
|
||||
b := foldComplement(got + seq + fl + l4)
|
||||
if a != b {
|
||||
t.Fatalf("dataOff=%d seq=%#x fl=%#x l4=%d: %#04x != %#04x",
|
||||
dataOff, seq, fl, l4, a, b)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -550,9 +550,7 @@ func (t *tun) Read(to []byte) (int, error) {
|
||||
return n - 4, nil
|
||||
}
|
||||
|
||||
// Write pushes one IP packet onto the utun device. Safe for concurrent use:
|
||||
// the AF prefix and iovecs are per-call stack state, and the fd write itself
|
||||
// serializes on the runtime's fd mutex (see the Queue contract in tio.go).
|
||||
// Write pushes one IP packet onto the utun device. Only valid for single threaded use.
|
||||
func (t *tun) Write(from []byte) (int, error) {
|
||||
if len(from) == 0 {
|
||||
return 0, syscall.EIO
|
||||
|
||||
+92
-43
@@ -25,17 +25,32 @@ import (
|
||||
)
|
||||
|
||||
type tun struct {
|
||||
readers tio.QueueSet
|
||||
closeLock sync.Mutex
|
||||
Device string
|
||||
vpnNetworks []netip.Prefix
|
||||
MaxMTU int
|
||||
DefaultMTU int
|
||||
TXQueueLen int
|
||||
deviceIndex int
|
||||
ioctlFd uintptr
|
||||
vnetHdr bool
|
||||
readers tio.QueueSet
|
||||
closeLock sync.Mutex
|
||||
Device string
|
||||
vpnNetworks []netip.Prefix
|
||||
MaxMTU int
|
||||
DefaultMTU int
|
||||
TXQueueLen int
|
||||
deviceIndex int
|
||||
ioctlFd uintptr
|
||||
vnetHdr bool
|
||||
// offloadFlags is the exact TUN_F_* offload mask newTun negotiated with
|
||||
// the kernel: usoOffloadFlags when USO was accepted, tsoOffloadFlags on
|
||||
// the TSO-only fallback, or 0 when vnetHdr is off. TUNSETOFFLOAD is
|
||||
// device-wide (drivers/net/tun.c set_offload updates tun->set_features
|
||||
// for the whole netdev), so addQueue must replay this exact
|
||||
// mask on every added queue — issuing a narrower mask there would
|
||||
// silently downgrade offloads (e.g. disable USO) for all queues while
|
||||
// they still advertise the stale capability.
|
||||
offloadFlags uint
|
||||
// routeFeatureECN, when true, sets RTAX_FEATURE_ECN on every route we
|
||||
// install for the tun. The kernel then actively negotiates ECN for
|
||||
// connections destined to those prefixes (equivalent to `ip route
|
||||
// change ... features ecn`) regardless of net.ipv4.tcp_ecn, so flows
|
||||
// across the nebula mesh use ECN even when the host default is the
|
||||
// passive setting (=2). Disable via tunnels.ecn=false.
|
||||
routeFeatureECN bool
|
||||
|
||||
Routes atomic.Pointer[[]Route]
|
||||
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
||||
@@ -76,7 +91,14 @@ type ifreqQLEN struct {
|
||||
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||
// We don't know what flags the caller opened this fd with and can't turn
|
||||
// on IFF_VNET_HDR after TUNSETIFF, so skip offload on inherited fds.
|
||||
return newTunGeneric(c, l, deviceFd, false, 0, vpnNetworks, "tun0")
|
||||
t, err := newTunGeneric(c, l, deviceFd, false, 0, vpnNetworks)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
t.Device = "tun0"
|
||||
|
||||
return t, nil
|
||||
}
|
||||
|
||||
// openTunDev opens /dev/net/tun, creating the device node first if it's
|
||||
@@ -102,7 +124,8 @@ func openTunDev() (int, error) {
|
||||
return fd, nil
|
||||
}
|
||||
|
||||
// tunSetIff runs TUNSETIFF with the given flags and returns the kernel-chosen device name on success.
|
||||
// tunSetIff runs TUNSETIFF with the given flags and returns the kernel-chosen
|
||||
// device name on success.
|
||||
func tunSetIff(fd int, name string, flags uint16) (string, error) {
|
||||
var req ifReq
|
||||
req.Flags = flags
|
||||
@@ -113,45 +136,57 @@ func tunSetIff(fd int, name string, flags uint16) (string, error) {
|
||||
return strings.Trim(string(req.Name[:]), "\x00"), nil
|
||||
}
|
||||
|
||||
// tsoOffloadFlags are the TUN_F_* bits we ask the kernel to enable when a TSO-capable TUN is available.
|
||||
// tsoOffloadFlags are the TUN_F_* bits we ask the kernel to enable when a
|
||||
// TSO-capable TUN is available. CSUM is required as a prerequisite for TSO.
|
||||
// TSO_ECN tells the kernel we propagate ECN correctly through coalesce and
|
||||
// segmentation, so it can deliver superpackets whose seed has CWR/ECE set
|
||||
// or whose IP-level codepoint is CE.
|
||||
const tsoOffloadFlags = unix.TUN_F_CSUM | unix.TUN_F_TSO4 | unix.TUN_F_TSO6 | unix.TUN_F_TSO_ECN
|
||||
|
||||
// usoAndTSOOffloadFlags adds UDP Segmentation Offload to tsoOffloadFlags.
|
||||
// Requires Linux >= 6.2; older kernels reject it and we fall back to TCP-only TSO
|
||||
const usoAndTSOOffloadFlags = tsoOffloadFlags | unix.TUN_F_USO4 | unix.TUN_F_USO6
|
||||
// usoOffloadFlags adds UDP Segmentation Offload to tsoOffloadFlags. Requires
|
||||
// Linux ≥ 6.2; older kernels reject it and we fall back to TCP-only TSO via
|
||||
// tsoOffloadFlags.
|
||||
const usoOffloadFlags = tsoOffloadFlags | unix.TUN_F_USO4 | unix.TUN_F_USO6
|
||||
|
||||
// offloadUSOEnabled reports whether the negotiated offload mask includes UDP
|
||||
// Segmentation Offload. It is the single source of truth for the usoEnabled
|
||||
// capability surfaced by each queue, so the mask stored on the tun and the USO
|
||||
// bit reported to coalescers can never drift apart.
|
||||
func offloadUSOEnabled(offloadFlags uint) bool {
|
||||
return offloadFlags&(unix.TUN_F_USO4|unix.TUN_F_USO6) != 0
|
||||
}
|
||||
|
||||
func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) {
|
||||
// IFF_TUN_EXCL prevents us from attaching to an already-running tun
|
||||
baseFlags := uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_TUN_EXCL)
|
||||
baseFlags := uint16(unix.IFF_TUN | unix.IFF_NO_PI)
|
||||
if multiqueue {
|
||||
baseFlags |= unix.IFF_MULTI_QUEUE
|
||||
}
|
||||
nameStr := c.GetString("tun.dev", "")
|
||||
|
||||
// First try to enable IFF_VNET_HDR via TUNSETIFF and negotiate TUN_F_*
|
||||
// offloads via TUNSETOFFLOAD so we can receive TSO/USO superpackets.
|
||||
// We try TSO+USO first, fall back to TSO-only on kernels without USO
|
||||
// (Linux < 6.2), and finally give up on virtio headers entirely and
|
||||
// reopen as a plain TUN if neither offload mask is accepted.
|
||||
fd, err := openTunDev()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
vnetHdr := true
|
||||
|
||||
// First try to enable IFF_VNET_HDR via TUNSETIFF and negotiate TUN_F_* offloads
|
||||
// We try TSO+USO first, fall back to TSO-only on kernels without USO (Linux < 6.2),
|
||||
// and finally give up on virtio headers entirely and reopen as a plain TUN if neither offload mask is accepted.
|
||||
|
||||
// offloadFlags is the exact TUN_F_* mask the kernel accepted.
|
||||
// We save it so addQueue can replay the identical device-wide mask on added queues
|
||||
// offloadFlags is the exact TUN_F_* mask the kernel accepted. We remember
|
||||
// it (rather than a plain bool) so addQueue can replay the
|
||||
// identical device-wide mask on added queues instead of downgrading them.
|
||||
var offloadFlags uint
|
||||
name, err := tunSetIff(fd, nameStr, baseFlags|unix.IFF_VNET_HDR)
|
||||
if err != nil {
|
||||
_ = unix.Close(fd)
|
||||
vnetHdr = false
|
||||
} else {
|
||||
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(usoAndTSOOffloadFlags)); err == nil {
|
||||
offloadFlags = usoAndTSOOffloadFlags
|
||||
// Try TSO+USO first. On kernels without USO support (Linux < 6.2)
|
||||
// the ioctl returns EINVAL; fall back to the TCP-only mask before
|
||||
// giving up on VNET_HDR entirely.
|
||||
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(usoOffloadFlags)); err == nil {
|
||||
offloadFlags = usoOffloadFlags
|
||||
} else if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(tsoOffloadFlags)); err == nil {
|
||||
offloadFlags = tsoOffloadFlags
|
||||
} else {
|
||||
@@ -177,7 +212,7 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
|
||||
l.Info("TUN offload enabled", "tso", true, "uso", offloadUSOEnabled(offloadFlags))
|
||||
}
|
||||
|
||||
t, err := newTunGeneric(c, l, fd, vnetHdr, offloadFlags, vpnNetworks, name)
|
||||
t, err := newTunGeneric(c, l, fd, vnetHdr, offloadFlags, vpnNetworks)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -187,14 +222,16 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
|
||||
return t, nil
|
||||
}
|
||||
|
||||
// newTunGeneric does all the stuff common to different tun initialization paths.
|
||||
// It will close your files on error.
|
||||
// offloadFlags is the TUN_F_* mask newTun negotiated (ignored when vnetHdr is false)
|
||||
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFlags uint, vpnNetworks []netip.Prefix, name string) (*tun, error) {
|
||||
// newTunGeneric does all the stuff common to different tun initialization
|
||||
// paths. It will close your files on error. offloadFlags is the TUN_F_* mask
|
||||
// newTun negotiated (0 when vnetHdr is off); the queues' USO capability is
|
||||
// derived from it so it can never disagree with the mask we replay on added
|
||||
// multiqueue readers.
|
||||
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFlags uint, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||
var qs tio.QueueSet
|
||||
var err error
|
||||
if vnetHdr {
|
||||
qs, err = tio.NewOffloadQueueSet(offloadUSOEnabled(offloadFlags), l)
|
||||
qs, err = tio.NewOffloadQueueSet(offloadUSOEnabled(offloadFlags))
|
||||
} else {
|
||||
qs, err = tio.NewPollQueueSet()
|
||||
}
|
||||
@@ -205,15 +242,11 @@ func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFla
|
||||
}
|
||||
err = qs.Add(fd)
|
||||
if err != nil {
|
||||
// Add only appends on success, so closing the set here can't
|
||||
// double-close fd; it releases the set's shutdown eventfd.
|
||||
_ = unix.Close(fd)
|
||||
_ = qs.Close()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
t := &tun{
|
||||
Device: name,
|
||||
readers: qs,
|
||||
closeLock: sync.Mutex{},
|
||||
vnetHdr: vnetHdr,
|
||||
@@ -222,6 +255,7 @@ func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFla
|
||||
TXQueueLen: c.GetInt("tun.tx_queue", 500),
|
||||
useSystemRoutes: c.GetBool("tun.use_system_route_table", false),
|
||||
useSystemRoutesBufferSize: c.GetInt("tun.use_system_route_table_buffer_size", 0),
|
||||
routeFeatureECN: c.GetBool("tunnels.ecn", true),
|
||||
routesFromSystem: map[netip.Prefix]routing.Gateways{},
|
||||
l: l,
|
||||
}
|
||||
@@ -315,7 +349,9 @@ func (t *tun) reload(c *config.C, initial bool) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Queues opens additional kernel multiqueue fds until the device has n queues, then returns them all.
|
||||
// Queues opens additional kernel multiqueue fds until the device has n
|
||||
// queues, then returns them all. The first queue was opened by newTun; each
|
||||
// extra fd replays the negotiated offload state (see addQueue).
|
||||
func (t *tun) Queues(n int) ([]tio.Queue, error) {
|
||||
for len(t.readers.Queues()) < n {
|
||||
if err := t.addQueue(); err != nil {
|
||||
@@ -325,7 +361,8 @@ func (t *tun) Queues(n int) ([]tio.Queue, error) {
|
||||
return t.readers.Queues(), nil
|
||||
}
|
||||
|
||||
// addQueue opens one more IFF_MULTI_QUEUE fd on the device and adds it to the queue set.
|
||||
// addQueue opens one more IFF_MULTI_QUEUE fd on the device and adds it to
|
||||
// the queue set.
|
||||
func (t *tun) addQueue() error {
|
||||
t.closeLock.Lock()
|
||||
defer t.closeLock.Unlock()
|
||||
@@ -345,6 +382,10 @@ func (t *tun) addQueue() error {
|
||||
}
|
||||
|
||||
if t.vnetHdr {
|
||||
// Replay the exact mask newTun negotiated. TUNSETOFFLOAD is
|
||||
// device-wide, so issuing the TSO-only mask here would disable USO
|
||||
// for every queue (including queue 0) on kernels where newTun
|
||||
// successfully enabled it, while the queues keep advertising USO.
|
||||
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(t.offloadFlags)); err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return fmt.Errorf("failed to enable offload on multiqueue tun fd: %w", err)
|
||||
@@ -525,13 +566,18 @@ func (t *tun) setDefaultRoute(cidr netip.Prefix) error {
|
||||
Table: unix.RT_TABLE_MAIN,
|
||||
Type: unix.RTN_UNICAST,
|
||||
}
|
||||
// Match the metric the kernel uses for its auto-installed connected route,
|
||||
// so RouteReplace overwrites it in place instead of adding a second route at a worse metric.
|
||||
// IPv6 connected routes are installed at metric 256 (IP6_RT_PRIO_KERN); IPv4 uses 0.
|
||||
// Without this, the kernel route wins lookups and our MTU / AdvMSS / Features never apply on v6.
|
||||
// Match the metric the kernel uses for its auto-installed connected
|
||||
// route, so RouteReplace overwrites it in place instead of adding a
|
||||
// second route at a worse metric. IPv6 connected routes are installed
|
||||
// at metric 256 (IP6_RT_PRIO_KERN); IPv4 uses 0. Without this, the
|
||||
// kernel route wins lookups and our MTU / AdvMSS / Features never
|
||||
// apply on v6.
|
||||
if cidr.Addr().Is6() {
|
||||
nr.Priority = 256
|
||||
}
|
||||
if t.routeFeatureECN {
|
||||
nr.Features |= unix.RTAX_FEATURE_ECN
|
||||
}
|
||||
err := netlink.RouteReplace(&nr)
|
||||
if err != nil {
|
||||
t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr)
|
||||
@@ -581,6 +627,9 @@ func (t *tun) addRoutes(logErrors bool) error {
|
||||
if r.Metric > 0 {
|
||||
nr.Priority = r.Metric
|
||||
}
|
||||
if t.routeFeatureECN {
|
||||
nr.Features |= unix.RTAX_FEATURE_ECN
|
||||
}
|
||||
|
||||
err := netlink.RouteReplace(&nr)
|
||||
if err != nil {
|
||||
|
||||
@@ -39,14 +39,14 @@ func TestTunAdvMSS(t *testing.T) {
|
||||
// capability: it is derived from the negotiated offload mask, so the mask
|
||||
// stored on the tun and the capability reported to coalescers cannot drift.
|
||||
func TestOffloadUSOEnabled(t *testing.T) {
|
||||
// usoAndTSOOffloadFlags must be a strict superset of tsoOffloadFlags. Otherwise
|
||||
// usoOffloadFlags must be a strict superset of tsoOffloadFlags. Otherwise
|
||||
// the TSO-only fallback (and the historic hardcoded-mask bug in
|
||||
// addQueue) would not actually be a downgrade.
|
||||
if usoAndTSOOffloadFlags&tsoOffloadFlags != tsoOffloadFlags {
|
||||
t.Fatalf("usoAndTSOOffloadFlags (%#x) is not a superset of tsoOffloadFlags (%#x)", usoAndTSOOffloadFlags, tsoOffloadFlags)
|
||||
if usoOffloadFlags&tsoOffloadFlags != tsoOffloadFlags {
|
||||
t.Fatalf("usoOffloadFlags (%#x) is not a superset of tsoOffloadFlags (%#x)", usoOffloadFlags, tsoOffloadFlags)
|
||||
}
|
||||
if usoAndTSOOffloadFlags == tsoOffloadFlags {
|
||||
t.Fatal("usoAndTSOOffloadFlags must add bits beyond tsoOffloadFlags")
|
||||
if usoOffloadFlags == tsoOffloadFlags {
|
||||
t.Fatal("usoOffloadFlags must add bits beyond tsoOffloadFlags")
|
||||
}
|
||||
|
||||
cases := []struct {
|
||||
@@ -54,7 +54,7 @@ func TestOffloadUSOEnabled(t *testing.T) {
|
||||
offloadFlags uint
|
||||
wantUSO bool
|
||||
}{
|
||||
{"uso-negotiated", usoAndTSOOffloadFlags, true},
|
||||
{"uso-negotiated", usoOffloadFlags, true},
|
||||
{"tso-fallback", tsoOffloadFlags, false},
|
||||
{"no-vnet-hdr", 0, false},
|
||||
}
|
||||
@@ -78,12 +78,12 @@ func TestOffloadUSOEnabled(t *testing.T) {
|
||||
// TUNSETOFFLOAD argument is read from.
|
||||
func TestAddQueueReplaysNegotiatedMask(t *testing.T) {
|
||||
t.Run("uso-negotiated", func(t *testing.T) {
|
||||
tn := &tun{vnetHdr: true, offloadFlags: usoAndTSOOffloadFlags}
|
||||
tn := &tun{vnetHdr: true, offloadFlags: usoOffloadFlags}
|
||||
// The ioctl argument in addQueue is uintptr(t.offloadFlags);
|
||||
// it must equal the negotiated USO mask, and must NOT be the TSO-only
|
||||
// mask (the original bug).
|
||||
if tn.offloadFlags != usoAndTSOOffloadFlags {
|
||||
t.Fatalf("offloadFlags = %#x, want %#x", tn.offloadFlags, usoAndTSOOffloadFlags)
|
||||
if tn.offloadFlags != usoOffloadFlags {
|
||||
t.Fatalf("offloadFlags = %#x, want %#x", tn.offloadFlags, usoOffloadFlags)
|
||||
}
|
||||
if tn.offloadFlags == tsoOffloadFlags {
|
||||
t.Fatal("added queue would downgrade USO: offloadFlags must not be the TSO-only mask when USO was negotiated")
|
||||
|
||||
+4
-6
@@ -10,13 +10,11 @@ import (
|
||||
_ "net/http/pprof" // registers pprof handlers on http.DefaultServeMux
|
||||
)
|
||||
|
||||
// startPprofServer serves net/http/pprof on localhost:6060 for the life of
|
||||
// ctx. It is only compiled into debug builds (`-tags debug`, `make debug`),
|
||||
// so a debug build announces itself with the Info line below. Loopback only:
|
||||
// a wildcard bind would expose profiles (peer addresses, config-derived
|
||||
// state) to anything that can reach the host, the overlay included.
|
||||
// startPprofServer serves net/http/pprof on :6060 for the life of ctx. It is
|
||||
// only compiled into debug builds (`-tags debug`, `make debug`), so a debug
|
||||
// build announces itself with the Info line below.
|
||||
func startPprofServer(ctx context.Context, l *slog.Logger) {
|
||||
server := &http.Server{Addr: "localhost:6060", Handler: nil}
|
||||
server := &http.Server{Addr: ":6060", Handler: nil}
|
||||
l.Info("Starting pprof debug server (debug build)", "addr", server.Addr)
|
||||
|
||||
go func() {
|
||||
|
||||
+1
-1
@@ -161,7 +161,7 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak
|
||||
switch existingRelay.State {
|
||||
case Established:
|
||||
hl.Log(context.Background(), level, "Send handshake via relay", "relay", relay.String())
|
||||
f.SendVia(relayHostInfo, existingRelay, stage0, make([]byte, 12), make([]byte, mtu), false, 0)
|
||||
f.SendVia(relayHostInfo, existingRelay, stage0, make([]byte, 12), make([]byte, mtu), false)
|
||||
case Disestablished:
|
||||
// Mark this relay as 'requested'
|
||||
relayHostInfo.relayState.UpdateRelayForByIpState(vpnIp, Requested)
|
||||
|
||||
+28
-13
@@ -14,28 +14,43 @@ const MTU = 9001
|
||||
// only costs additional sendmmsg chunks within a single WriteBatch call.
|
||||
const MaxWriteBatch = 128
|
||||
|
||||
// RxMeta carries per-packet metadata extracted from the RX path (ancillary
|
||||
// data, kernel offload state, etc.) and passed to EncReader callbacks.
|
||||
// Backends that do not produce a particular signal leave its zero value.
|
||||
//
|
||||
// OuterECN is the 2-bit IP-level ECN codepoint stamped on the carrier
|
||||
// datagram (extracted from IP_TOS / IPV6_TCLASS cmsg on Linux). Zero
|
||||
// means Not-ECT, which is also the value backends without ECN RX support
|
||||
// supply on every packet.
|
||||
type RxMeta struct {
|
||||
OuterECN byte
|
||||
}
|
||||
|
||||
type EncReader func(
|
||||
addr netip.AddrPort,
|
||||
payload []byte,
|
||||
meta RxMeta,
|
||||
)
|
||||
|
||||
type Conn interface {
|
||||
Rebind() error
|
||||
LocalAddr() (netip.AddrPort, error)
|
||||
// ListenOut invokes r for each received packet.
|
||||
// On batch-capable backends (recvmmsg), flush is called after each batch is fully delivered.
|
||||
// Callers use it to flush per-batch accumulators such as TUN write coalescers.
|
||||
// Single-packet backends call flush after each packet. flush must not be nil.
|
||||
// ListenOut invokes r for each received packet. On batch-capable
|
||||
// backends (recvmmsg), flush is called after each batch is fully
|
||||
// delivered — callers use it to flush per-batch accumulators such as
|
||||
// TUN write coalescers. Single-packet backends call flush after each
|
||||
// packet. flush must not be nil.
|
||||
ListenOut(r EncReader, flush func()) error
|
||||
WriteTo(b []byte, addr netip.AddrPort) error
|
||||
// WriteBatch sends a contiguous batch of packets, each with its own
|
||||
// destination. bufs and addrs must have the same length. Linux uses
|
||||
// sendmmsg(2) for a single syscall.
|
||||
//
|
||||
// Returns the number of packets successfully written. A destination the kernel rejects costs only
|
||||
// its own packet, so a short count means some peers were undeliverable, not that the batch failed.
|
||||
// Not safe for concurrent use on the same Conn.
|
||||
WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error)
|
||||
// destination. bufs and addrs must have the same length. outerECNs may
|
||||
// be nil (treated as all-zero / Not-ECT); when non-nil it must have the
|
||||
// same length as bufs, and outerECNs[i] is the 2-bit IP-level ECN
|
||||
// codepoint to set on packet i's outer header. Linux uses sendmmsg(2)
|
||||
// for a single syscall and attaches the value as IP_TOS / IPV6_TCLASS
|
||||
// cmsg; other backends ignore it. Returns on the first error; callers
|
||||
// may observe a partial send if some packets went out before the error.
|
||||
WriteBatch(bufs [][]byte, addrs []netip.AddrPort, outerECNs []byte) error
|
||||
ReloadConfig(c *config.C)
|
||||
SupportsMultipleReaders() bool
|
||||
Close() error
|
||||
@@ -58,8 +73,8 @@ func (NoopConn) SupportsMultipleReaders() bool {
|
||||
func (NoopConn) WriteTo(_ []byte, _ netip.AddrPort) error {
|
||||
return nil
|
||||
}
|
||||
func (NoopConn) WriteBatch(bufs [][]byte, _ []netip.AddrPort) (int, error) {
|
||||
return len(bufs), nil
|
||||
func (NoopConn) WriteBatch(_ [][]byte, _ []netip.AddrPort, _ []byte) error {
|
||||
return nil
|
||||
}
|
||||
func (NoopConn) ReloadConfig(_ *config.C) {
|
||||
return
|
||||
|
||||
@@ -1,61 +0,0 @@
|
||||
package udp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
)
|
||||
|
||||
// NetworkChangeMonitor rebinds the udp listener when the local network moves out from under it.
|
||||
//
|
||||
// Detection lives here in the udp package, next to the socket it concerns and the platform matrix that already knows
|
||||
// which sockets go stale. What to do about a change — updating the lighthouse, requerying tunnels — is not the udp
|
||||
// package's business, so Start takes the reaction as a plain function. Passing it at Start rather than holding it
|
||||
// keeps this package from referencing whatever owns the rebind.
|
||||
//
|
||||
// On platforms whose sockets do not go stale, watchNetworkChanges hands back a nil channel and Start returns.
|
||||
type NetworkChangeMonitor struct {
|
||||
l *slog.Logger
|
||||
ctx context.Context
|
||||
enabled bool
|
||||
}
|
||||
|
||||
// NewNetworkChangeMonitor builds a monitor for local network changes. The returned monitor is always usable: Start
|
||||
// is safe to call unconditionally, it no-ops when disabled or on a platform that does not need it.
|
||||
func NewNetworkChangeMonitor(ctx context.Context, l *slog.Logger, c *config.C) *NetworkChangeMonitor {
|
||||
return &NetworkChangeMonitor{
|
||||
l: l,
|
||||
ctx: ctx,
|
||||
enabled: c.GetBool("listen.rebind_on_network_change", true),
|
||||
}
|
||||
}
|
||||
|
||||
// Start watches for network changes until the context is cancelled, calling rebind once per settled change. It
|
||||
// blocks, so callers run it in a goroutine, and it no-ops when disabled, unsupported, or with nothing to rebind.
|
||||
func (m *NetworkChangeMonitor) Start(rebind func()) {
|
||||
if !m.enabled || rebind == nil || m.ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
|
||||
changes, err := watchNetworkChanges(m.ctx, m.l)
|
||||
if err != nil {
|
||||
// Not fatal. Everything else still works, we just won't notice a network change on our own.
|
||||
m.l.Error("Failed to watch for network changes, will not rebind the udp listener when the network moves",
|
||||
"error", err,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
if changes == nil {
|
||||
// This platform's sockets don't go stale, so there is nothing to watch for.
|
||||
return
|
||||
}
|
||||
|
||||
m.l.Info("Watching for network changes to rebind the udp listener")
|
||||
|
||||
for range changes {
|
||||
m.l.Info("Local network changed, rebinding the udp listener")
|
||||
rebind()
|
||||
}
|
||||
}
|
||||
@@ -1,164 +0,0 @@
|
||||
//go:build darwin && !ios && !e2e_testing
|
||||
// +build darwin,!ios,!e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
const (
|
||||
// netChangeSettleWindow is how long we keep swallowing routing messages after the first interesting one. A
|
||||
// single network change is never a single message, it is a burst: the link drops, addresses go away, new ones
|
||||
// arrive, routes get rewritten. Reporting part way through that just means reporting again.
|
||||
netChangeSettleWindow = time.Second
|
||||
|
||||
// netChangeReadBuffer is sized well past any rt_msghdr plus its addresses. A short read would be discarded by
|
||||
// the kernel, so being generous here is how we avoid missing a message.
|
||||
netChangeReadBuffer = 4096
|
||||
)
|
||||
|
||||
// watchNetworkChanges reports when the local network moves out from under us, so the listener can be rebound.
|
||||
//
|
||||
// Darwin scopes a udp socket to whatever interface it came up on. Move between networks and we keep sending out an
|
||||
// interface that no longer has a route, which surfaces as an instant "no route to host" with no packet ever leaving
|
||||
// the box. Rebind clears that, but only if something notices the change and calls it. iOS has always been told by
|
||||
// the host app off NWPathMonitor. This is the equivalent for everything else that runs on darwin.
|
||||
//
|
||||
// The returned channel is buffered and coalescing: a send is dropped if one is already pending, since both mean the
|
||||
// same thing to a reader. It is closed when ctx is cancelled or the routing socket fails, so a caller can simply
|
||||
// range over it. Platforms whose sockets do not need rebinding return a nil channel and no error.
|
||||
func watchNetworkChanges(ctx context.Context, l *slog.Logger) (<-chan struct{}, error) {
|
||||
sock, err := openRouteSocket()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
changes := make(chan struct{}, 1)
|
||||
|
||||
go func() {
|
||||
defer close(changes)
|
||||
defer func() { _ = sock.Close() }()
|
||||
|
||||
// Closing the socket is what unblocks the read in watchRouteSocket, so this turns cancellation into a
|
||||
// close. It is scoped to this call so it cannot outlive the watch it belongs to.
|
||||
done := make(chan struct{})
|
||||
defer close(done)
|
||||
go func() {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
_ = sock.Close()
|
||||
case <-done:
|
||||
}
|
||||
}()
|
||||
|
||||
watchRouteSocket(l, sock, changes)
|
||||
}()
|
||||
|
||||
return changes, nil
|
||||
}
|
||||
|
||||
// watchRouteSocket blocks reading the routing socket, reporting once per settled burst of changes. It returns when
|
||||
// the socket is closed, which is how cancellation gets us out of here.
|
||||
func watchRouteSocket(l *slog.Logger, sock *os.File, changes chan<- struct{}) {
|
||||
buf := make([]byte, netChangeReadBuffer)
|
||||
|
||||
for {
|
||||
n, err := sock.Read(buf)
|
||||
if err != nil {
|
||||
logRouteSocketError(l, err)
|
||||
return
|
||||
}
|
||||
|
||||
if !isNetworkChange(buf[:n]) {
|
||||
continue
|
||||
}
|
||||
|
||||
// Swallow the rest of the burst. The deadline is absolute and not extended by what arrives, so this always
|
||||
// ends after the settle window no matter how chatty the socket is. Changes that land after the window
|
||||
// simply produce another report, which is the correct outcome anyway.
|
||||
deadline := time.Now().Add(netChangeSettleWindow)
|
||||
for {
|
||||
if err = sock.SetReadDeadline(deadline); err != nil {
|
||||
logRouteSocketError(l, err)
|
||||
return
|
||||
}
|
||||
|
||||
if _, err = sock.Read(buf); err != nil {
|
||||
if os.IsTimeout(err) {
|
||||
break
|
||||
}
|
||||
logRouteSocketError(l, err)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if err = sock.SetReadDeadline(time.Time{}); err != nil {
|
||||
logRouteSocketError(l, err)
|
||||
return
|
||||
}
|
||||
|
||||
select {
|
||||
case changes <- struct{}{}:
|
||||
default:
|
||||
// One already pending, and a second "the network moved" tells the reader nothing new.
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// logRouteSocketError reports a routing socket failure unless it is just us shutting the socket down.
|
||||
func logRouteSocketError(l *slog.Logger, err error) {
|
||||
if errors.Is(err, os.ErrClosed) {
|
||||
return
|
||||
}
|
||||
|
||||
l.Error("Error reading the routing socket, will no longer notice local network changes", "error", err)
|
||||
}
|
||||
|
||||
// openRouteSocket returns the routing socket as a non blocking os.File. Going through os.File puts reads on the go
|
||||
// poller, which buys us both a working read deadline and a Close that unblocks a read in progress.
|
||||
func openRouteSocket() (*os.File, error) {
|
||||
fd, err := unix.Socket(unix.AF_ROUTE, unix.SOCK_RAW, unix.AF_UNSPEC)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err = unix.SetNonblock(fd, true); err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return os.NewFile(uintptr(fd), "route"), nil
|
||||
}
|
||||
|
||||
// isNetworkChange reports whether a routing message means our local addressing may have moved out from under us.
|
||||
//
|
||||
// We read the header instead of parsing the message because the type is the only part we need, and a full parse can
|
||||
// fail on shapes we don't care about, which would turn "a message I can't parse" into "a change I missed".
|
||||
// rt_msghdr, if_msghdr and ifa_msghdr all begin with the same three fields, so this is the same for every type.
|
||||
func isNetworkChange(msg []byte) bool {
|
||||
if len(msg) < 4 {
|
||||
return false
|
||||
}
|
||||
|
||||
// u_short msglen, u_char version, u_char type
|
||||
if int(binary.NativeEndian.Uint16(msg[0:2])) > len(msg) || msg[2] != unix.RTM_VERSION {
|
||||
return false
|
||||
}
|
||||
|
||||
switch msg[3] {
|
||||
case unix.RTM_NEWADDR, unix.RTM_DELADDR, unix.RTM_IFINFO:
|
||||
// An address arrived or left, or a link changed state. Anything else on this socket is either a route
|
||||
// churning underneath us, which a rebind doesn't help with, or unrelated traffic.
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
@@ -1,244 +0,0 @@
|
||||
//go:build darwin && !ios && !e2e_testing
|
||||
// +build darwin,!ios,!e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/goleak"
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// routeMsg builds the first four bytes of a routing message, which is all isNetworkChange reads.
|
||||
func routeMsg(msgType uint8, extra int) []byte {
|
||||
msg := make([]byte, 4+extra)
|
||||
binary.NativeEndian.PutUint16(msg[0:2], uint16(len(msg)))
|
||||
msg[2] = unix.RTM_VERSION
|
||||
msg[3] = msgType
|
||||
return msg
|
||||
}
|
||||
|
||||
func TestIsNetworkChange(t *testing.T) {
|
||||
// The three that mean our addressing may have moved
|
||||
assert.True(t, isNetworkChange(routeMsg(unix.RTM_NEWADDR, 0)))
|
||||
assert.True(t, isNetworkChange(routeMsg(unix.RTM_DELADDR, 0)))
|
||||
assert.True(t, isNetworkChange(routeMsg(unix.RTM_IFINFO, 0)))
|
||||
|
||||
// Route churn is not something a rebind helps with
|
||||
assert.False(t, isNetworkChange(routeMsg(unix.RTM_ADD, 0)))
|
||||
assert.False(t, isNetworkChange(routeMsg(unix.RTM_DELETE, 0)))
|
||||
assert.False(t, isNetworkChange(routeMsg(unix.RTM_GET, 0)))
|
||||
|
||||
// Garbage must not be mistaken for a change
|
||||
assert.False(t, isNetworkChange(nil), "empty")
|
||||
assert.False(t, isNetworkChange([]byte{0, 0, 0}), "short header")
|
||||
|
||||
wrongVersion := routeMsg(unix.RTM_NEWADDR, 0)
|
||||
wrongVersion[2] = unix.RTM_VERSION + 1
|
||||
assert.False(t, isNetworkChange(wrongVersion), "wrong rtm_version")
|
||||
|
||||
lying := routeMsg(unix.RTM_NEWADDR, 0)
|
||||
binary.NativeEndian.PutUint16(lying[0:2], 512)
|
||||
assert.False(t, isNetworkChange(lying), "msglen longer than what we read")
|
||||
}
|
||||
|
||||
// socketPair returns a connected pair of datagram sockets, the first wrapped the same way the routing socket is. It
|
||||
// stands in for the kernel so the watch loop can be driven with synthetic messages.
|
||||
func socketPair(t *testing.T) (*os.File, int) {
|
||||
t.Helper()
|
||||
|
||||
fds, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_DGRAM, 0)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, unix.SetNonblock(fds[0], true))
|
||||
|
||||
f := os.NewFile(uintptr(fds[0]), "route")
|
||||
t.Cleanup(func() {
|
||||
_ = f.Close()
|
||||
_ = unix.Close(fds[1])
|
||||
})
|
||||
|
||||
return f, fds[1]
|
||||
}
|
||||
|
||||
func TestWatchRouteSocketCoalescesABurst(t *testing.T) {
|
||||
sock, kernel := socketPair(t)
|
||||
changes := make(chan struct{}, 1)
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
watchRouteSocket(test.NewLogger(), sock, changes)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// One network change is a burst of messages. All of these land inside the settle window, so they must produce
|
||||
// exactly one report rather than one apiece.
|
||||
for range 5 {
|
||||
_, err := unix.Write(kernel, routeMsg(unix.RTM_NEWADDR, 8))
|
||||
require.NoError(t, err)
|
||||
}
|
||||
// Uninteresting messages in the middle of a burst must not add a report of their own either.
|
||||
_, err := unix.Write(kernel, routeMsg(unix.RTM_ADD, 8))
|
||||
require.NoError(t, err)
|
||||
|
||||
select {
|
||||
case <-changes:
|
||||
case <-time.After(netChangeSettleWindow * 4):
|
||||
t.Fatal("a burst should have reported a change")
|
||||
}
|
||||
|
||||
// Nothing more from that burst
|
||||
select {
|
||||
case <-changes:
|
||||
t.Fatal("a burst should report exactly once")
|
||||
case <-time.After(netChangeSettleWindow):
|
||||
}
|
||||
|
||||
// A change after the window has closed is a separate event and gets its own report.
|
||||
_, err = unix.Write(kernel, routeMsg(unix.RTM_IFINFO, 8))
|
||||
require.NoError(t, err)
|
||||
select {
|
||||
case <-changes:
|
||||
case <-time.After(netChangeSettleWindow * 4):
|
||||
t.Fatal("a later change should report again")
|
||||
}
|
||||
|
||||
// Closing the socket is how the real thing shuts down
|
||||
require.NoError(t, sock.Close())
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("watchRouteSocket did not return after the socket was closed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWatchRouteSocketIgnoresUninterestingMessages(t *testing.T) {
|
||||
sock, kernel := socketPair(t)
|
||||
changes := make(chan struct{}, 1)
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
watchRouteSocket(test.NewLogger(), sock, changes)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
for _, msgType := range []uint8{unix.RTM_ADD, unix.RTM_DELETE, unix.RTM_GET, unix.RTM_MISS} {
|
||||
_, err := unix.Write(kernel, routeMsg(msgType, 8))
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
select {
|
||||
case <-changes:
|
||||
t.Fatal("route churn alone must not report a change")
|
||||
case <-time.After(netChangeSettleWindow * 2):
|
||||
}
|
||||
|
||||
require.NoError(t, sock.Close())
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("watchRouteSocket did not return after the socket was closed")
|
||||
}
|
||||
}
|
||||
|
||||
// TestWatchRouteSocketDropsRatherThanBlocks covers the coalescing send. A reader that is busy rebinding must not
|
||||
// wedge the watcher, and a second pending "the network moved" tells it nothing new anyway.
|
||||
func TestWatchRouteSocketDropsRatherThanBlocks(t *testing.T) {
|
||||
sock, kernel := socketPair(t)
|
||||
changes := make(chan struct{}, 1)
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
watchRouteSocket(test.NewLogger(), sock, changes)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// Nobody is reading changes, so after the first report the buffer is full for the rest of this test
|
||||
for range 3 {
|
||||
_, err := unix.Write(kernel, routeMsg(unix.RTM_NEWADDR, 8))
|
||||
require.NoError(t, err)
|
||||
time.Sleep(netChangeSettleWindow + time.Millisecond*250)
|
||||
}
|
||||
|
||||
// The watcher must still be alive and responsive to a close
|
||||
require.NoError(t, sock.Close())
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("watchRouteSocket wedged on a full channel")
|
||||
}
|
||||
|
||||
assert.Len(t, changes, 1, "the pending report should have coalesced, not queued")
|
||||
}
|
||||
|
||||
// TestWatchNetworkChangesStopsWithContext covers the detection path against a real routing socket, including that
|
||||
// cancelling the context closes the channel so a ranging caller falls out of its loop.
|
||||
func TestWatchNetworkChangesStopsWithContext(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
changes, err := watchNetworkChanges(ctx, test.NewLogger())
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, changes, "darwin should support watching")
|
||||
|
||||
drained := make(chan struct{})
|
||||
go func() {
|
||||
for range changes {
|
||||
}
|
||||
close(drained)
|
||||
}()
|
||||
|
||||
cancel()
|
||||
select {
|
||||
case <-drained:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("cancelling the context should close the changes channel")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNetworkChangeMonitorStopsWithContext drives the whole monitor against a real routing socket: Start must block
|
||||
// watching, and cancelling the context (which is all Control does on shutdown, it never stops the monitor directly)
|
||||
// must return it and clean up the watch goroutines.
|
||||
func TestNetworkChangeMonitorStopsWithContext(t *testing.T) {
|
||||
// IgnoreCurrent because other tests in this package leave readers running; we only care about what this test
|
||||
// leaks itself.
|
||||
defer goleak.VerifyNone(t, goleak.IgnoreCurrent())
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
l := test.NewLogger()
|
||||
c := config.NewC(l)
|
||||
require.NoError(t, c.LoadString("listen:\n rebind_on_network_change: true\n"))
|
||||
m := NewNetworkChangeMonitor(ctx, l, c)
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
m.Start(func() {})
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// Start should be sitting on the routing socket, not have fallen out. If it returned early it either failed to
|
||||
// watch or no-op'd, both of which we want to catch.
|
||||
select {
|
||||
case <-done:
|
||||
t.Fatal("Start returned instead of watching")
|
||||
case <-time.After(time.Millisecond * 250):
|
||||
}
|
||||
|
||||
cancel()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("Start did not return after the context was cancelled")
|
||||
}
|
||||
|
||||
// Starting again after the context is dead must not open anything.
|
||||
m.Start(func() {})
|
||||
}
|
||||
@@ -1,22 +0,0 @@
|
||||
//go:build !darwin || ios || e2e_testing
|
||||
// +build !darwin ios e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
)
|
||||
|
||||
// watchNetworkChanges is a no-op outside of darwin.
|
||||
//
|
||||
// Darwin is the platform that scopes a udp socket to the interface it came up on, so it is the platform whose socket
|
||||
// goes stale when the local network changes. Everywhere else Rebind has nothing to do, so there is nothing to watch
|
||||
// for. iOS is excluded on purpose even though it is darwin: the host app already drives the rebind off NWPathMonitor,
|
||||
// and two things racing to rebind the same socket is worse than one.
|
||||
//
|
||||
// A nil channel means "not supported here", which callers must treat as "do not start a watcher" rather than
|
||||
// selecting on it, since a receive from a nil channel blocks forever.
|
||||
func watchNetworkChanges(_ context.Context, _ *slog.Logger) (<-chan struct{}, error) {
|
||||
return nil, nil
|
||||
}
|
||||
@@ -1,39 +0,0 @@
|
||||
package udp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func newMonitor(t *testing.T, ctx context.Context, cfg string) *NetworkChangeMonitor {
|
||||
t.Helper()
|
||||
l := test.NewLogger()
|
||||
c := config.NewC(l)
|
||||
require.NoError(t, c.LoadString(cfg))
|
||||
return NewNetworkChangeMonitor(ctx, l, c)
|
||||
}
|
||||
|
||||
func TestNetworkChangeMonitorDefaultsOn(t *testing.T) {
|
||||
// Says nothing about rebinding, so this covers the default.
|
||||
m := newMonitor(t, context.Background(), "listen:\n host: 0.0.0.0\n")
|
||||
assert.True(t, m.enabled, "should default to on")
|
||||
}
|
||||
|
||||
func TestNetworkChangeMonitorDisabledIsANoOp(t *testing.T) {
|
||||
m := newMonitor(t, context.Background(), "listen:\n rebind_on_network_change: false\n")
|
||||
require.False(t, m.enabled)
|
||||
|
||||
// Must return without opening a socket. If it watched anything this would block.
|
||||
m.Start(func() {})
|
||||
}
|
||||
|
||||
func TestNetworkChangeMonitorNilRebindIsANoOp(t *testing.T) {
|
||||
// Nothing to rebind, so there is no point watching, on any platform.
|
||||
m := newMonitor(t, context.Background(), "listen:\n rebind_on_network_change: true\n")
|
||||
m.Start(nil)
|
||||
}
|
||||
@@ -0,0 +1,62 @@
|
||||
//go:build !android && !e2e_testing
|
||||
// +build !android,!e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
import (
|
||||
"net"
|
||||
"syscall"
|
||||
"unsafe"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// rawSendmmsg performs sendmmsg(2) over a syscall.RawConn without
|
||||
// allocating a closure per call. The struct holds preallocated in/out
|
||||
// scratch (chunk/sent/errno) and a method-value bound at construction so
|
||||
// rawConn.Write receives a stable function pointer instead of a fresh
|
||||
// closure on every send.
|
||||
type rawSendmmsg struct {
|
||||
msgs []rawMessage
|
||||
chunk int
|
||||
sent int
|
||||
errno syscall.Errno
|
||||
callback func(fd uintptr) bool
|
||||
}
|
||||
|
||||
// bind wires r.callback to r.run. Must be called once after r.msgs is set;
|
||||
// subsequent send calls invoke r.callback without rebinding.
|
||||
func (r *rawSendmmsg) bind() { r.callback = r.run }
|
||||
|
||||
// run is the preallocated callback rawConn.Write invokes. It reads its
|
||||
// input (r.chunk) and writes its outputs (r.sent, r.errno) through the
|
||||
// rawSendmmsg fields so the method value does not capture per-call locals
|
||||
// and therefore does not heap-allocate.
|
||||
func (r *rawSendmmsg) run(fd uintptr) bool {
|
||||
r1, _, errno := unix.Syscall6(unix.SYS_SENDMMSG, fd,
|
||||
uintptr(unsafe.Pointer(&r.msgs[0])), uintptr(r.chunk),
|
||||
0, 0, 0,
|
||||
)
|
||||
if errno == syscall.EAGAIN || errno == syscall.EWOULDBLOCK {
|
||||
return false
|
||||
}
|
||||
r.sent = int(r1)
|
||||
r.errno = errno
|
||||
return true
|
||||
}
|
||||
|
||||
// send issues sendmmsg over rc against the first n entries of r.msgs.
|
||||
// Returns the number of entries the kernel processed and any error;
|
||||
// matches the original sendmmsg helper's contract.
|
||||
func (r *rawSendmmsg) send(rc syscall.RawConn, n int) (int, error) {
|
||||
r.chunk = n
|
||||
r.sent = 0
|
||||
r.errno = 0
|
||||
if err := rc.Write(r.callback); err != nil {
|
||||
return r.sent, err
|
||||
}
|
||||
if r.errno != 0 {
|
||||
return r.sent, &net.OpError{Op: "sendmmsg", Err: r.errno}
|
||||
}
|
||||
return r.sent, nil
|
||||
}
|
||||
+10
-16
@@ -140,20 +140,13 @@ func (u *StdConn) WriteTo(b []byte, ap netip.AddrPort) error {
|
||||
}
|
||||
}
|
||||
|
||||
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
|
||||
// An un-sendable destination costs its own packet, never the ones behind it in the batch.
|
||||
// TODO: WriteTo maps EWOULDBLOCK to an error, so a full send buffer
|
||||
// silently drops the rest of a burst (linux blocks instead). Poll for
|
||||
// writability on EAGAIN before giving up on the remainder.
|
||||
written := 0
|
||||
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||
for i, b := range bufs {
|
||||
if err := u.WriteTo(b, addrs[i]); err == nil {
|
||||
written++
|
||||
} else {
|
||||
u.l.Debug("failed to write packet in batch", "udpAddr", addrs[i], "error", err)
|
||||
if err := u.WriteTo(b, addrs[i]); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return written, nil
|
||||
return nil
|
||||
}
|
||||
|
||||
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
||||
@@ -195,7 +188,7 @@ func (u *StdConn) ListenOut(r EncReader, flush func()) error {
|
||||
continue
|
||||
}
|
||||
|
||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n:n])
|
||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n], RxMeta{})
|
||||
flush()
|
||||
}
|
||||
}
|
||||
@@ -204,9 +197,6 @@ func (u *StdConn) SupportsMultipleReaders() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// Rebind clears the interface the kernel scoped this socket to, so that sends are routed against the current
|
||||
// routing table instead of the interface we happened to be on when the socket was created. Darwin pins sockets
|
||||
// this way on its own, which is what strands us after the underlying network changes.
|
||||
func (u *StdConn) Rebind() error {
|
||||
var err error
|
||||
if u.isV4 {
|
||||
@@ -215,5 +205,9 @@ func (u *StdConn) Rebind() error {
|
||||
err = syscall.SetsockoptInt(int(u.sysFd), syscall.IPPROTO_IPV6, syscall.IPV6_BOUND_IF, 0)
|
||||
}
|
||||
|
||||
return err
|
||||
if err != nil {
|
||||
u.l.Error("Failed to rebind udp socket", "error", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// TestPlanRunBreaksOnECNChange confirms that two same-destination, same-size
|
||||
// packets with different outer ECN end up in separate sendmmsg entries (the
|
||||
// kernel stamps one outer codepoint per entry, so a run that straddled the
|
||||
// boundary would silently lose information).
|
||||
func TestPlanRunBreaksOnECNChange(t *testing.T) {
|
||||
u := &StdConn{gsoSupported: true, maxGSOSegments: 63}
|
||||
dst := netip.MustParseAddrPort("10.0.0.1:4242")
|
||||
|
||||
bufs := [][]byte{
|
||||
make([]byte, 1200),
|
||||
make([]byte, 1200),
|
||||
make([]byte, 1200),
|
||||
}
|
||||
addrs := []netip.AddrPort{dst, dst, dst}
|
||||
|
||||
t.Run("uniform_ecn_runs_together", func(t *testing.T) {
|
||||
ecns := []byte{0x02, 0x02, 0x02}
|
||||
runLen, segSize := u.planRun(bufs, addrs, ecns, 0, 64)
|
||||
if runLen != 3 {
|
||||
t.Errorf("runLen=%d want 3 (uniform ECT(0))", runLen)
|
||||
}
|
||||
if segSize != 1200 {
|
||||
t.Errorf("segSize=%d want 1200", segSize)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("ecn_change_truncates_run", func(t *testing.T) {
|
||||
// 0,0,3: first two run together, CE seeds a fresh entry.
|
||||
ecns := []byte{0x00, 0x00, 0x03}
|
||||
runLen, _ := u.planRun(bufs, addrs, ecns, 0, 64)
|
||||
if runLen != 2 {
|
||||
t.Errorf("runLen=%d want 2 (ECN changes at index 2)", runLen)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("nil_ecns_runs_full", func(t *testing.T) {
|
||||
runLen, _ := u.planRun(bufs, addrs, nil, 0, 64)
|
||||
if runLen != 3 {
|
||||
t.Errorf("runLen=%d want 3 (nil ecns means no break)", runLen)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("first_ecn_is_singleton", func(t *testing.T) {
|
||||
// Second packet has different ECN from the first → run halts at 1
|
||||
// (the first packet alone forms the run).
|
||||
ecns := []byte{0x00, 0x03, 0x03}
|
||||
runLen, _ := u.planRun(bufs, addrs, ecns, 0, 64)
|
||||
if runLen != 1 {
|
||||
t.Errorf("runLen=%d want 1 (different ECN immediately)", runLen)
|
||||
}
|
||||
})
|
||||
}
|
||||
+5
-9
@@ -44,17 +44,13 @@ func (u *GenericConn) WriteTo(b []byte, addr netip.AddrPort) error {
|
||||
return err
|
||||
}
|
||||
|
||||
func (u *GenericConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
|
||||
// An un-sendable destination costs its own packet, never the ones behind it in the batch.
|
||||
written := 0
|
||||
func (u *GenericConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||
for i, b := range bufs {
|
||||
if _, err := u.UDPConn.WriteToUDPAddrPort(b, addrs[i]); err == nil {
|
||||
written++
|
||||
} else {
|
||||
u.l.Debug("failed to write packet in batch", "udpAddr", addrs[i], "error", err)
|
||||
if _, err := u.UDPConn.WriteToUDPAddrPort(b, addrs[i]); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return written, nil
|
||||
return nil
|
||||
}
|
||||
|
||||
func (u *GenericConn) LocalAddr() (netip.AddrPort, error) {
|
||||
@@ -106,7 +102,7 @@ func (u *GenericConn) ListenOut(r EncReader, flush func()) error {
|
||||
continue
|
||||
}
|
||||
|
||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n:n])
|
||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n], RxMeta{})
|
||||
flush()
|
||||
}
|
||||
}
|
||||
|
||||
+716
-196
File diff suppressed because it is too large
Load Diff
@@ -30,6 +30,39 @@ type rawMessage struct {
|
||||
Len uint32
|
||||
}
|
||||
|
||||
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
|
||||
msgs := make([]rawMessage, n)
|
||||
buffers := make([][]byte, n)
|
||||
names := make([][]byte, n)
|
||||
|
||||
var cmsgs []byte
|
||||
if cmsgSpace > 0 {
|
||||
cmsgs = make([]byte, n*cmsgSpace)
|
||||
}
|
||||
|
||||
for i := range msgs {
|
||||
buffers[i] = make([]byte, bufSize)
|
||||
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||
|
||||
vs := []iovec{
|
||||
{Base: &buffers[i][0], Len: uint32(len(buffers[i]))},
|
||||
}
|
||||
|
||||
msgs[i].Hdr.Iov = &vs[0]
|
||||
msgs[i].Hdr.Iovlen = uint32(len(vs))
|
||||
|
||||
msgs[i].Hdr.Name = &names[i][0]
|
||||
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
||||
|
||||
if cmsgSpace > 0 {
|
||||
msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace]
|
||||
msgs[i].Hdr.Controllen = uint32(cmsgSpace)
|
||||
}
|
||||
}
|
||||
|
||||
return msgs, buffers, names, cmsgs
|
||||
}
|
||||
|
||||
func setIovLen(v *iovec, n int) {
|
||||
v.Len = uint32(n)
|
||||
}
|
||||
|
||||
@@ -33,6 +33,39 @@ type rawMessage struct {
|
||||
Pad0 [4]byte
|
||||
}
|
||||
|
||||
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
|
||||
msgs := make([]rawMessage, n)
|
||||
buffers := make([][]byte, n)
|
||||
names := make([][]byte, n)
|
||||
|
||||
var cmsgs []byte
|
||||
if cmsgSpace > 0 {
|
||||
cmsgs = make([]byte, n*cmsgSpace)
|
||||
}
|
||||
|
||||
for i := range msgs {
|
||||
buffers[i] = make([]byte, bufSize)
|
||||
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||
|
||||
vs := []iovec{
|
||||
{Base: &buffers[i][0], Len: uint64(len(buffers[i]))},
|
||||
}
|
||||
|
||||
msgs[i].Hdr.Iov = &vs[0]
|
||||
msgs[i].Hdr.Iovlen = uint64(len(vs))
|
||||
|
||||
msgs[i].Hdr.Name = &names[i][0]
|
||||
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
||||
|
||||
if cmsgSpace > 0 {
|
||||
msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace]
|
||||
msgs[i].Hdr.Controllen = uint64(cmsgSpace)
|
||||
}
|
||||
}
|
||||
|
||||
return msgs, buffers, names, cmsgs
|
||||
}
|
||||
|
||||
func setIovLen(v *iovec, n int) {
|
||||
v.Len = uint64(n)
|
||||
}
|
||||
|
||||
+100
-555
@@ -3,11 +3,11 @@
|
||||
package udp
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"encoding/binary"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"syscall"
|
||||
"testing"
|
||||
"time"
|
||||
"unsafe"
|
||||
@@ -54,6 +54,39 @@ func buildCmsg(level, typ int32, data []byte) []byte {
|
||||
return buf
|
||||
}
|
||||
|
||||
// TestParseRecvCmsgOuterECNFamily is the RX half of the dual-stack ECN fix:
|
||||
// parseRecvCmsg must read the outer ECN from whichever family the kernel
|
||||
// delivered, not from the socket family. On the default `::` dual-stack bind
|
||||
// a v4 peer's outer ECN arrives as an IP_TOS cmsg, which the old socket-family
|
||||
// gate ignored entirely.
|
||||
func TestParseRecvCmsgOuterECNFamily(t *testing.T) {
|
||||
tc := make([]byte, 4)
|
||||
binary.NativeEndian.PutUint32(tc, 0x02)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
ctrl []byte
|
||||
want byte
|
||||
}{
|
||||
{"ip_tos_ce", buildCmsg(int32(unix.IPPROTO_IP), int32(unix.IP_TOS), []byte{0x03}), 0x03},
|
||||
{"ip_tos_ect0", buildCmsg(int32(unix.IPPROTO_IP), int32(unix.IP_TOS), []byte{0x02}), 0x02},
|
||||
{"ipv6_tclass_ect0", buildCmsg(int32(unix.IPPROTO_IPV6), int32(unix.IPV6_TCLASS), tc), 0x02},
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
hdr := &msghdr{Control: &c.ctrl[0]}
|
||||
setMsgControllen(hdr, len(c.ctrl))
|
||||
gso, ecn := parseRecvCmsg(hdr, false, true)
|
||||
if gso != 0 {
|
||||
t.Errorf("gso = %d, want 0 (no UDP_GRO cmsg present)", gso)
|
||||
}
|
||||
if ecn != c.want {
|
||||
t.Errorf("ecn = 0x%02x, want 0x%02x", ecn, c.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func testLogger() *slog.Logger {
|
||||
return slog.New(slog.DiscardHandler)
|
||||
}
|
||||
@@ -90,13 +123,9 @@ func TestWriteBatchBadFamilyDeliversOthers(t *testing.T) {
|
||||
bufs := [][]byte{[]byte("AAA"), []byte("BBB"), []byte("CCC")}
|
||||
addrs := []netip.AddrPort{good, bad, good}
|
||||
|
||||
n, err := sender.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
if err := sender.WriteBatch(bufs, addrs, nil); err != nil {
|
||||
t.Fatalf("WriteBatch returned error, want nil (bad dest should be isolated): %v", err)
|
||||
}
|
||||
if n != 2 {
|
||||
t.Errorf("WriteBatch wrote %d packets, want 2 of 3 (the bad-family dest is the only casualty)", n)
|
||||
}
|
||||
|
||||
got := map[string]bool{}
|
||||
rx.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
@@ -116,11 +145,12 @@ func TestWriteBatchBadFamilyDeliversOthers(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteBatchUnreachableDestDeliversOthers is the kernel-rejection twin of
|
||||
// TestWriteBatchBadFamilyDeliversOthers. A destination the kernel refuses outright (240.0.0.0/4 is reserved, so
|
||||
// the send returns EINVAL) fails its sendmmsg entry; WriteBatch must drop only that entry and still deliver
|
||||
// every other packet rather than abandoning the batch at the first failure.
|
||||
func TestWriteBatchUnreachableDestDeliversOthers(t *testing.T) {
|
||||
// TestWriteBatchOuterTOSToV4Mapped is the TX half of the dual-stack ECN fix,
|
||||
// verified against a live kernel: WriteBatch on the default `::` dual-stack
|
||||
// socket, sending to a v4-mapped destination, must stamp the outer ECN via an
|
||||
// IP_TOS cmsg (not IPV6_TCLASS, which the kernel's v4 path ignores) so a v4
|
||||
// receiver actually sees it.
|
||||
func TestWriteBatchOuterTOSToV4Mapped(t *testing.T) {
|
||||
rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
|
||||
if err != nil {
|
||||
t.Skipf("cannot open v4 receiver (sandbox?): %v", err)
|
||||
@@ -128,561 +158,76 @@ func TestWriteBatchUnreachableDestDeliversOthers(t *testing.T) {
|
||||
defer rx.Close()
|
||||
rxPort := rx.LocalAddr().(*net.UDPAddr).Port
|
||||
|
||||
c, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, false, 1)
|
||||
// Ask the kernel to deliver the received outer TOS as ancillary data.
|
||||
rxRaw, err := rx.SyscallConn()
|
||||
if err != nil {
|
||||
t.Skipf("cannot open v4 sender (sandbox?): %v", err)
|
||||
t.Fatalf("SyscallConn: %v", err)
|
||||
}
|
||||
var soErr error
|
||||
if err := rxRaw.Control(func(fd uintptr) {
|
||||
soErr = unix.SetsockoptInt(int(fd), unix.IPPROTO_IP, unix.IP_RECVTOS, 1)
|
||||
}); err != nil || soErr != nil {
|
||||
t.Skipf("cannot enable IP_RECVTOS (sandbox/kernel?): ctrl=%v so=%v", err, soErr)
|
||||
}
|
||||
|
||||
c, err := NewListener(testLogger(), netip.IPv6Unspecified(), 0, false, 1)
|
||||
if err != nil {
|
||||
t.Skipf("cannot open dual-stack sender (sandbox?): %v", err)
|
||||
}
|
||||
defer c.Close()
|
||||
sender := c.(*StdConn)
|
||||
|
||||
good := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), uint16(rxPort))
|
||||
bad := netip.MustParseAddrPort("240.0.0.1:9999") // reserved space, the kernel refuses it
|
||||
|
||||
bufs := [][]byte{[]byte("P0"), []byte("P1"), []byte("BAD"), []byte("P3"), []byte("P4")}
|
||||
addrs := []netip.AddrPort{good, good, bad, good, good}
|
||||
|
||||
// The bad destination is reported, but only after every other packet has been attempted.
|
||||
if _, err := sender.WriteBatch(bufs, addrs); err == nil {
|
||||
t.Log("WriteBatch returned nil; kernel accepted the reserved address, delivery assertions still apply")
|
||||
if sender.isV4 {
|
||||
t.Skipf("sender came up v4-only; need a dual-stack v6 socket for this test")
|
||||
}
|
||||
|
||||
got := map[string]bool{}
|
||||
rx.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
buf := make([]byte, 64)
|
||||
for i := 0; i < 4; i++ {
|
||||
n, _, rerr := rx.ReadFromUDPAddrPort(buf)
|
||||
if rerr != nil {
|
||||
t.Fatalf("expected 4 delivered packets, read #%d failed: %v (got so far: %v)", i+1, rerr, got)
|
||||
}
|
||||
got[string(buf[:n])] = true
|
||||
}
|
||||
for _, want := range []string{"P0", "P1", "P3", "P4"} {
|
||||
if !got[want] {
|
||||
t.Errorf("packet %s was not delivered; delivered set = %v", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
// v4-mapped-in-v6 destination: routed through the kernel's IPv4 path.
|
||||
dst := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), uint16(rxPort))
|
||||
const wantECN = byte(0x02) // ECT(0)
|
||||
|
||||
// TestParseRecvCmsgCorruptLenNoPanic: a cmsg Len near max-int used to wrap
|
||||
// off+clen negative, slip past the bounds check, and drive the walk offset
|
||||
// negative -- a panic on the next ctrl[off]. The guard must compare Len
|
||||
// against the remaining bytes instead. Also pins the plain truncated-Len
|
||||
// cases (too small, larger than the buffer) to a clean early return.
|
||||
func TestParseRecvCmsgCorruptLenNoPanic(t *testing.T) {
|
||||
// First cmsg: a valid empty one so the walk advances past off=0
|
||||
// (off+clen can't overflow while off is still zero).
|
||||
valid := buildCmsg(int32(unix.SOL_UDP), int32(unix.UDP_GRO), make([]byte, 4))
|
||||
|
||||
corrupt := func(lenVal int) []byte {
|
||||
buf := make([]byte, len(valid)+unix.CmsgSpace(4))
|
||||
copy(buf, valid)
|
||||
h := (*unix.Cmsghdr)(unsafe.Pointer(&buf[len(valid)]))
|
||||
h.Level = int32(unix.IPPROTO_IP)
|
||||
h.Type = int32(unix.IP_TOS)
|
||||
setCmsgLen(h, lenVal)
|
||||
return buf
|
||||
}
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
ctrl []byte
|
||||
}{
|
||||
{"len_near_max_int", corrupt(int(^uint(0)>>1) - 8)},
|
||||
{"len_too_small", corrupt(unix.SizeofCmsghdr - 1)},
|
||||
{"len_past_buffer", corrupt(1 << 20)},
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
hdr := &msghdr{Control: &c.ctrl[0]}
|
||||
setMsgControllen(hdr, len(c.ctrl))
|
||||
gso := parseRecvCmsg(hdr)
|
||||
// The valid leading UDP_GRO cmsg (payload 0) must still parse;
|
||||
// the corrupt trailer just ends the walk.
|
||||
if gso != 0 {
|
||||
t.Errorf("parseRecvCmsg = %d, want 0", gso)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestDeliverSegments pins the GRO RX splitting: a kernel-coalesced buffer
|
||||
// must come back out as the exact pre-coalesce packets -- every boundary
|
||||
// error here shreds encrypted packets and every decrypt downstream fails.
|
||||
func TestDeliverSegments(t *testing.T) {
|
||||
from := netip.MustParseAddrPort("192.0.2.1:4242")
|
||||
// Spare backing capacity mimics the recvmmsg row a real payload sits in;
|
||||
// the cap checks below prove none of it leaks to a delivered segment.
|
||||
pay := func(n int) []byte {
|
||||
b := make([]byte, n, n+512)
|
||||
for i := range b {
|
||||
b[i] = byte(i)
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
payload []byte
|
||||
segSize int
|
||||
wantLens []int
|
||||
}{
|
||||
{"no-gro", pay(1400), 0, []int{1400}},
|
||||
{"negative-segsize", pay(1400), -5, []int{1400}},
|
||||
{"segsize-equals-payload", pay(1400), 1400, []int{1400}},
|
||||
{"segsize-past-payload", pay(1400), 2000, []int{1400}},
|
||||
{"even-split", pay(4200), 1400, []int{1400, 1400, 1400}},
|
||||
{"short-tail", pay(3000), 1400, []int{1400, 1400, 200}},
|
||||
{"single-byte-tail", pay(2801), 1400, []int{1400, 1400, 1}},
|
||||
{"segsize-one", pay(3), 1, []int{1, 1, 1}},
|
||||
{"empty-payload", pay(0), 1400, []int{0}},
|
||||
{"max-coalesce", pay(65500), 1372, nil}, // lens derived below
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
wantLens := c.wantLens
|
||||
if wantLens == nil {
|
||||
for rem := len(c.payload); rem > 0; rem -= c.segSize {
|
||||
wantLens = append(wantLens, min(c.segSize, rem))
|
||||
}
|
||||
}
|
||||
|
||||
var got [][]byte
|
||||
deliverSegments(func(a netip.AddrPort, seg []byte) {
|
||||
if a != from {
|
||||
t.Errorf("from = %v, want %v", a, from)
|
||||
}
|
||||
got = append(got, seg)
|
||||
}, from, c.payload, c.segSize)
|
||||
|
||||
if len(got) != len(wantLens) {
|
||||
t.Fatalf("delivered %d segments, want %d", len(got), len(wantLens))
|
||||
}
|
||||
// Segments must tile the payload in order with no gap, overlap,
|
||||
// or copy: each must alias the payload at the right offset.
|
||||
off := 0
|
||||
for i, seg := range got {
|
||||
if len(seg) != wantLens[i] {
|
||||
t.Fatalf("segment %d len=%d want %d", i, len(seg), wantLens[i])
|
||||
}
|
||||
if cap(seg) != len(seg) {
|
||||
// EncReader contract: an append into spare capacity would
|
||||
// scribble into the next segment of the shared row.
|
||||
t.Errorf("segment %d cap=%d, want %d (capacity must not reach into the row)", i, cap(seg), len(seg))
|
||||
}
|
||||
if len(seg) > 0 && &seg[0] != &c.payload[off] {
|
||||
t.Errorf("segment %d does not alias payload at offset %d", i, off)
|
||||
}
|
||||
off += len(seg)
|
||||
}
|
||||
if off != len(c.payload) {
|
||||
t.Errorf("segments cover %d bytes, payload has %d", off, len(c.payload))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// newRewindTestWriter builds a batchWriter with no socket: GSO planning on,
|
||||
// sendFn left for the test to script. fd is invalid on purpose -- any path
|
||||
// that actually hits the kernel fails loudly.
|
||||
func newRewindTestWriter() *batchWriter {
|
||||
w := &batchWriter{fd: -1, isV4: true, l: testLogger()}
|
||||
w.prepareWriteMessages(MaxWriteBatch)
|
||||
w.gsoSupported = true
|
||||
w.maxGSOSegments = 63
|
||||
return w
|
||||
}
|
||||
|
||||
// capturePrepared decodes n prepared mmsghdr entries beginning at start
|
||||
// straight from their iovecs -- ground truth, deliberately not the entryEnd
|
||||
// bookkeeping the resume logic itself relies on. Returns one []byte per
|
||||
// packed packet, in entry order.
|
||||
func capturePrepared(w *batchWriter, start, n int) [][]byte {
|
||||
var out [][]byte
|
||||
for e := start; e < start+n; e++ {
|
||||
hdr := &w.msgs[e].Hdr
|
||||
iovs := unsafe.Slice(hdr.Iov, int(hdr.Iovlen))
|
||||
for _, iov := range iovs {
|
||||
b := make([]byte, int(iov.Len))
|
||||
if iov.Len > 0 {
|
||||
copy(b, unsafe.Slice(iov.Base, int(iov.Len)))
|
||||
}
|
||||
out = append(out, b)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// TestWriteBatchPartialSendRewind drives WriteBatch through scripted
|
||||
// partial sendmmsg results and asserts the rewind resumes exactly where
|
||||
// the kernel stopped: every packet on the wire exactly once, in order,
|
||||
// no duplicate, no loss. This is the hairiest logic in the write path
|
||||
// and a rewind bug means silent packet duplication or loss under EAGAIN-
|
||||
// style backpressure.
|
||||
func TestWriteBatchPartialSendRewind(t *testing.T) {
|
||||
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
|
||||
|
||||
mkBuf := func(tag byte, n int) []byte {
|
||||
b := make([]byte, n)
|
||||
for i := range b {
|
||||
b[i] = tag
|
||||
}
|
||||
b[0] = tag // tag identifies the packet uniquely below
|
||||
return b
|
||||
}
|
||||
|
||||
// Mixed shape: a 3-packet GSO run to A, a lone short packet to A (run
|
||||
// tail), then two to B. The planner packs this as multiple entries with
|
||||
// multi-iovec runs, which is what makes the rewind arithmetic hairy.
|
||||
bufs := [][]byte{
|
||||
mkBuf(1, 1200), mkBuf(2, 1200), mkBuf(3, 1200), // run to A
|
||||
mkBuf(4, 600), // short tail to A
|
||||
mkBuf(5, 900), mkBuf(6, 900), // run to B
|
||||
}
|
||||
addrs := []netip.AddrPort{dstA, dstA, dstA, dstA, dstB, dstB}
|
||||
|
||||
scripts := [][]int{
|
||||
{99}, // accept everything first call
|
||||
{1, 99}, // one entry per call, then the rest
|
||||
{1, 1, 1, 99}, // strictly one entry per call
|
||||
{2, 99}, // two entries, then the rest
|
||||
}
|
||||
for si, script := range scripts {
|
||||
t.Run(fmt.Sprintf("script_%d", si), func(t *testing.T) {
|
||||
w := newRewindTestWriter()
|
||||
var wire [][]byte
|
||||
call := 0
|
||||
w.sendFn = func(start, n int) (int, error) {
|
||||
accept := n
|
||||
if call < len(script) && script[call] < n {
|
||||
accept = script[call]
|
||||
}
|
||||
call++
|
||||
wire = append(wire, capturePrepared(w, start, accept)...)
|
||||
return accept, nil
|
||||
}
|
||||
|
||||
written, err := w.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteBatch: %v", err)
|
||||
}
|
||||
if written != len(bufs) {
|
||||
t.Errorf("written = %d, want %d", written, len(bufs))
|
||||
}
|
||||
if len(wire) != len(bufs) {
|
||||
t.Fatalf("wire got %d packets, want %d (dup or loss in rewind)", len(wire), len(bufs))
|
||||
}
|
||||
for i, b := range wire {
|
||||
if len(b) != len(bufs[i]) || b[0] != bufs[i][0] {
|
||||
t.Errorf("wire[%d] = tag %d len %d, want tag %d len %d (reorder/dup)",
|
||||
i, b[0], len(b), bufs[i][0], len(bufs[i]))
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteBatchSkipUnroutableRunAccounting: an unroutable destination mid-
|
||||
// batch is skipped without committing an entry, leaving a hole in the bufs
|
||||
// index space. The written count must tally packets per sent entry -- the
|
||||
// index span would count the hole -- across both full and partial sendmmsg
|
||||
// success.
|
||||
func TestWriteBatchSkipUnroutableRunAccounting(t *testing.T) {
|
||||
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
|
||||
bad := netip.MustParseAddrPort("[2001:db8::1]:9999") // v6 dest, v4 writer
|
||||
|
||||
mk := func(tag byte, n int) []byte {
|
||||
b := make([]byte, n)
|
||||
b[0] = tag
|
||||
return b
|
||||
}
|
||||
bufs := [][]byte{mk(1, 1200), mk(2, 1200), mk(3, 500), mk(4, 900), mk(5, 900)}
|
||||
addrs := []netip.AddrPort{dstA, dstA, bad, dstB, dstB}
|
||||
|
||||
for si, script := range [][]int{{99}, {1, 99}} {
|
||||
t.Run(fmt.Sprintf("script_%d", si), func(t *testing.T) {
|
||||
w := newRewindTestWriter()
|
||||
var wire [][]byte
|
||||
call := 0
|
||||
w.sendFn = func(start, n int) (int, error) {
|
||||
accept := n
|
||||
if call < len(script) && script[call] < n {
|
||||
accept = script[call]
|
||||
}
|
||||
call++
|
||||
wire = append(wire, capturePrepared(w, start, accept)...)
|
||||
return accept, nil
|
||||
}
|
||||
|
||||
written, err := w.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteBatch: %v", err)
|
||||
}
|
||||
if written != 4 {
|
||||
t.Errorf("written = %d, want 4 (the unroutable run is the only casualty)", written)
|
||||
}
|
||||
wantTags := []byte{1, 2, 4, 5}
|
||||
if len(wire) != len(wantTags) {
|
||||
t.Fatalf("wire got %d packets, want %d (dup or loss around the skip)", len(wire), len(wantTags))
|
||||
}
|
||||
for i, b := range wire {
|
||||
if b[0] != wantTags[i] {
|
||||
t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i])
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteBatchMidChunkRejectResumes: after a partial success, a zero-sent
|
||||
// error on the FIRST REMAINING entry (done > 0) must drop only that entry's
|
||||
// run and resume the rest of the chunk in place -- no repacking, no packets
|
||||
// lost from entries before or after the rejected one.
|
||||
func TestWriteBatchMidChunkRejectResumes(t *testing.T) {
|
||||
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
|
||||
dstC := netip.MustParseAddrPort("127.0.0.3:4242")
|
||||
|
||||
mk := func(tag byte, n int) []byte {
|
||||
b := make([]byte, n)
|
||||
b[0] = tag
|
||||
return b
|
||||
}
|
||||
// Three entries: a 2-packet GSO run to A, a 2-packet run to B, one to C.
|
||||
bufs := [][]byte{mk(1, 1200), mk(2, 1200), mk(3, 900), mk(4, 900), mk(5, 600)}
|
||||
addrs := []netip.AddrPort{dstA, dstA, dstB, dstB, dstC}
|
||||
|
||||
w := newRewindTestWriter()
|
||||
var wire [][]byte
|
||||
var starts []int
|
||||
call := 0
|
||||
w.sendFn = func(start, n int) (int, error) {
|
||||
starts = append(starts, start)
|
||||
call++
|
||||
switch call {
|
||||
case 1: // accept only entry 0 (the run to A)
|
||||
wire = append(wire, capturePrepared(w, start, 1)...)
|
||||
return 1, nil
|
||||
case 2: // reject entry 1 (the run to B) outright
|
||||
return -1, &net.OpError{Op: "sendmmsg", Err: unix.EPERM}
|
||||
default: // accept the rest
|
||||
wire = append(wire, capturePrepared(w, start, n)...)
|
||||
return n, nil
|
||||
}
|
||||
}
|
||||
|
||||
written, err := w.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
if err := sender.WriteBatch([][]byte{[]byte("tos-probe")}, []netip.AddrPort{dst}, []byte{wantECN}); err != nil {
|
||||
t.Fatalf("WriteBatch: %v", err)
|
||||
}
|
||||
if written != 3 {
|
||||
t.Errorf("written = %d, want 3 (B's rejected run is the only casualty)", written)
|
||||
|
||||
// Read the datagram plus its ancillary TOS.
|
||||
rx.SetReadDeadline(time.Now().Add(3 * time.Second))
|
||||
payload := make([]byte, 128)
|
||||
oob := make([]byte, 512)
|
||||
var n, oobn int
|
||||
var rerr error
|
||||
if err := rxRaw.Read(func(fd uintptr) bool {
|
||||
n, oobn, _, _, rerr = unix.Recvmsg(int(fd), payload, oob, 0)
|
||||
if rerr == syscall.EAGAIN || rerr == syscall.EWOULDBLOCK {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}); err != nil {
|
||||
t.Fatalf("waiting for datagram failed (no delivery?): %v", err)
|
||||
}
|
||||
wantTags := []byte{1, 2, 5}
|
||||
if len(wire) != len(wantTags) {
|
||||
t.Fatalf("wire got %d packets, want %d (dup or loss around the mid-chunk reject)", len(wire), len(wantTags))
|
||||
if rerr != nil {
|
||||
t.Fatalf("Recvmsg: %v", rerr)
|
||||
}
|
||||
for i, b := range wire {
|
||||
if b[0] != wantTags[i] {
|
||||
t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i])
|
||||
if string(payload[:n]) != "tos-probe" {
|
||||
t.Fatalf("payload = %q, want %q", string(payload[:n]), "tos-probe")
|
||||
}
|
||||
|
||||
cmsgs, err := unix.ParseSocketControlMessage(oob[:oobn])
|
||||
if err != nil {
|
||||
t.Fatalf("ParseSocketControlMessage: %v", err)
|
||||
}
|
||||
found := false
|
||||
var gotTOS byte
|
||||
for _, m := range cmsgs {
|
||||
if m.Header.Level == unix.IPPROTO_IP && m.Header.Type == unix.IP_TOS && len(m.Data) >= 1 {
|
||||
found = true
|
||||
gotTOS = m.Data[0]
|
||||
}
|
||||
}
|
||||
// The resume must reuse the prepared entries: same chunk, advancing
|
||||
// start offsets, no repack (which would restart at 0 with fresh entries).
|
||||
if want := []int{0, 1, 2}; !slices.Equal(starts, want) {
|
||||
t.Errorf("sendFn start offsets = %v, want %v", starts, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteBatchMidChunkEIODisablesGSOWithoutDup: an EIO on a GSO entry
|
||||
// after earlier entries in the chunk already went out must replay ONLY from
|
||||
// the failed run (replanned as single-packet entries) -- the already-sent
|
||||
// entries must not be duplicated.
|
||||
func TestWriteBatchMidChunkEIODisablesGSOWithoutDup(t *testing.T) {
|
||||
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
|
||||
|
||||
mk := func(tag byte, n int) []byte {
|
||||
b := make([]byte, n)
|
||||
b[0] = tag
|
||||
return b
|
||||
}
|
||||
// Entry 0: single packet to A. Entry 1: 2-packet GSO run to B.
|
||||
bufs := [][]byte{mk(1, 600), mk(2, 1200), mk(3, 1200)}
|
||||
addrs := []netip.AddrPort{dstA, dstB, dstB}
|
||||
|
||||
w := newRewindTestWriter()
|
||||
var wire [][]byte
|
||||
call := 0
|
||||
w.sendFn = func(start, n int) (int, error) {
|
||||
call++
|
||||
switch call {
|
||||
case 1: // accept entry 0 only
|
||||
wire = append(wire, capturePrepared(w, start, 1)...)
|
||||
return 1, nil
|
||||
case 2: // EIO on the GSO run to B
|
||||
return -1, &net.OpError{Op: "sendmmsg", Err: unix.EIO}
|
||||
default: // replanned single-packet replay
|
||||
wire = append(wire, capturePrepared(w, start, n)...)
|
||||
return n, nil
|
||||
}
|
||||
}
|
||||
|
||||
written, err := w.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteBatch: %v", err)
|
||||
}
|
||||
if w.gsoSupported {
|
||||
t.Error("gsoSupported still true after EIO on a GSO entry")
|
||||
}
|
||||
if written != len(bufs) {
|
||||
t.Errorf("written = %d, want %d", written, len(bufs))
|
||||
}
|
||||
wantTags := []byte{1, 2, 3}
|
||||
if len(wire) != len(wantTags) {
|
||||
t.Fatalf("wire got %d packets, want %d (packet 1 duplicated, or B's run lost)", len(wire), len(wantTags))
|
||||
}
|
||||
for i, b := range wire {
|
||||
if b[0] != wantTags[i] {
|
||||
t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteBatchZeroProgress: sent == 0 with no error must abort with an
|
||||
// error rather than spin forever replaying the same chunk.
|
||||
func TestWriteBatchZeroProgress(t *testing.T) {
|
||||
w := newRewindTestWriter()
|
||||
w.sendFn = func(start, n int) (int, error) { return 0, nil }
|
||||
bufs := [][]byte{make([]byte, 100)}
|
||||
addrs := []netip.AddrPort{netip.MustParseAddrPort("127.0.0.1:4242")}
|
||||
if _, err := w.WriteBatch(bufs, addrs); err == nil {
|
||||
t.Fatal("WriteBatch = nil error on zero progress, want error")
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteBatchEIODisablesGSOAndReplays pins the runtime GSO give-up: a
|
||||
// sendmmsg rejected with EIO on a GSO superpacket entry must clear
|
||||
// gsoSupported and replay the same packets as per-packet entries through
|
||||
// sendmmsg (keeping batching), not fall back to per-packet sendto.
|
||||
func TestWriteBatchEIODisablesGSOAndReplays(t *testing.T) {
|
||||
dst := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||
bufs := [][]byte{make([]byte, 1200), make([]byte, 1200), make([]byte, 1200)}
|
||||
addrs := []netip.AddrPort{dst, dst, dst}
|
||||
|
||||
w := newRewindTestWriter()
|
||||
var entryCounts []int
|
||||
call := 0
|
||||
w.sendFn = func(start, n int) (int, error) {
|
||||
entryCounts = append(entryCounts, n)
|
||||
call++
|
||||
if call == 1 {
|
||||
return -1, &net.OpError{Op: "sendmmsg", Err: unix.EIO}
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
written, err := w.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteBatch: %v", err)
|
||||
}
|
||||
if w.gsoSupported {
|
||||
t.Error("gsoSupported still true after EIO on a GSO entry")
|
||||
}
|
||||
if written != len(bufs) {
|
||||
t.Errorf("written = %d, want %d", written, len(bufs))
|
||||
}
|
||||
// First call: one GSO entry carrying the whole run. Replay: one entry
|
||||
// per packet, still via sendmmsg.
|
||||
want := []int{1, 3}
|
||||
if len(entryCounts) != len(want) || entryCounts[0] != want[0] || entryCounts[1] != want[1] {
|
||||
t.Errorf("sendmmsg entry counts = %v, want %v", entryCounts, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestGSOEngagesOnLoopback is the offload smoke test: real sockets, real
|
||||
// UDP_SEGMENT cmsg, real kernel segmentation over loopback. It asserts
|
||||
// both that GSO *engaged* (the whole batch left in a single sendmmsg
|
||||
// entry -- a silent fallback to per-packet entries fails the test) and
|
||||
// that the kernel carved the superpacket back into the exact original
|
||||
// datagrams on the receive side. Runs in CI (make test on ubuntu-latest),
|
||||
// which is what guards against the offload path silently degrading.
|
||||
func TestGSOEngagesOnLoopback(t *testing.T) {
|
||||
rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
|
||||
if err != nil {
|
||||
t.Fatalf("listen rx: %v", err)
|
||||
}
|
||||
defer rx.Close()
|
||||
dst := rx.LocalAddr().(*net.UDPAddr).AddrPort()
|
||||
|
||||
uc, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, false, 8)
|
||||
if err != nil {
|
||||
t.Fatalf("NewListener: %v", err)
|
||||
}
|
||||
sc := uc.(*StdConn)
|
||||
defer sc.Close()
|
||||
|
||||
if !sc.bw.gsoSupported {
|
||||
var un unix.Utsname
|
||||
_ = unix.Uname(&un)
|
||||
release := string(un.Release[:])
|
||||
if major, minor := parseRelease(release); major > 4 || (major == 4 && minor >= 18) {
|
||||
t.Fatalf("kernel %q supports UDP_SEGMENT but the GSO probe failed", release)
|
||||
}
|
||||
t.Skipf("kernel %q predates UDP_SEGMENT (4.18)", release)
|
||||
}
|
||||
|
||||
// Spy on the real syscall to count entries per sendmmsg without
|
||||
// changing what hits the kernel.
|
||||
var entryCounts []int
|
||||
real := sc.bw.sendFn
|
||||
sc.bw.sendFn = func(start, n int) (int, error) {
|
||||
entryCounts = append(entryCounts, n)
|
||||
return real(start, n)
|
||||
}
|
||||
|
||||
const numPkts = 8
|
||||
const pktLen = 1200
|
||||
bufs := make([][]byte, numPkts)
|
||||
addrs := make([]netip.AddrPort, numPkts)
|
||||
for i := range bufs {
|
||||
bufs[i] = make([]byte, pktLen)
|
||||
for j := range bufs[i] {
|
||||
bufs[i][j] = byte(i)
|
||||
}
|
||||
addrs[i] = dst
|
||||
}
|
||||
|
||||
written, err := sc.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteBatch: %v", err)
|
||||
}
|
||||
if written != numPkts {
|
||||
t.Fatalf("written = %d, want %d", written, numPkts)
|
||||
}
|
||||
// GSO engaged means the run went out as ONE sendmmsg entry carrying a
|
||||
// UDP_SEGMENT superpacket. Per-packet entries mean it silently fell
|
||||
// back -- exactly the regression this test exists to catch.
|
||||
if len(entryCounts) != 1 || entryCounts[0] != 1 {
|
||||
t.Fatalf("sendmmsg entry counts = %v, want [1]: GSO did not engage", entryCounts)
|
||||
}
|
||||
|
||||
// The kernel must deliver the original datagram boundaries and bytes.
|
||||
_ = rx.SetReadDeadline(time.Now().Add(5 * time.Second))
|
||||
got := make([]byte, pktLen+1)
|
||||
for i := 0; i < numPkts; i++ {
|
||||
n, _, err := rx.ReadFromUDP(got)
|
||||
if err != nil {
|
||||
t.Fatalf("rx read %d: %v", i, err)
|
||||
}
|
||||
if n != pktLen {
|
||||
t.Fatalf("rx read %d: len=%d want %d (kernel segmented at wrong boundary)", i, n, pktLen)
|
||||
}
|
||||
for j := 0; j < n; j++ {
|
||||
if got[j] != byte(i) {
|
||||
t.Fatalf("rx read %d: byte %d = %#x, want %#x", i, j, got[j], byte(i))
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatalf("no IP_TOS cmsg delivered to v4 receiver — outer ECN did not land (%d cmsgs)", len(cmsgs))
|
||||
}
|
||||
if gotTOS&0x03 != wantECN {
|
||||
t.Errorf("received outer TOS = 0x%02x, want low-2-bits = 0x%02x", gotTOS, wantECN)
|
||||
} else {
|
||||
t.Logf("verified: v4 receiver saw outer TOS 0x%02x (ECN=0x%02x) from dual-stack sender", gotTOS, gotTOS&0x03)
|
||||
}
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user