Compare commits

..

2 Commits

Author SHA1 Message Date
JackDoan bc04261d0b linux: let the kernel resolve the tun.dev %d template atomically
Resolving the template in userspace (LinkList then TUNSETIFF) races with
concurrent instances: both can pick nebula0, and with multiqueue the loser
silently attaches as a second queue to the winner's device instead of
failing. TUNSETIFF natively substitutes a single %d via dev_alloc_name,
atomically and always allocating a fresh device, so keep the fail-fast
validation but pass the template through and read the resolved name back.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Fg7oD3uZ4jZjjQ4gRsPAwB
2026-07-09 18:37:30 -05:00
JackDoan 57e1a9b6af linux: allow %d anywhere in the tun.dev template 2026-07-08 13:43:43 -05:00
138 changed files with 1452 additions and 15414 deletions
+2 -5
View File
@@ -25,9 +25,9 @@ inputs:
required: false
default: "code-signer"
key-prefix:
description: "S3 key prefix to write under; defaults to code-signing/<owner>/<repo> of the calling repo"
description: "S3 key prefix the caller is authorized to write under"
required: false
default: ""
default: "code-signing/slackhq/nebula"
runs:
using: composite
@@ -57,9 +57,6 @@ runs:
KEY_PREFIX: ${{ inputs.key-prefix }}
run: |
set -eu
# Default the prefix to this repo so the S3 key attributes the sign correctly.
# nebula-nightly runs this same action but writes under its own repo's prefix.
KEY_PREFIX="${KEY_PREFIX:-code-signing/$GITHUB_REPOSITORY}"
RUN="${GITHUB_RUN_ID}-${GITHUB_RUN_ATTEMPT}"
find "$SIGN_PATH" -name '*.exe' -print | while read -r path
+6 -6
View File
@@ -12,9 +12,9 @@ jobs:
steps:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: Build
@@ -38,9 +38,9 @@ jobs:
steps:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: Build
@@ -78,9 +78,9 @@ jobs:
steps:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: Import certificates
+6 -6
View File
@@ -32,9 +32,9 @@ jobs:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: add hashicorp source
@@ -64,9 +64,9 @@ jobs:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: add hashicorp source
@@ -90,9 +90,9 @@ jobs:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
# WSL2 + Ubuntu so the smoke can run a real linux peer with its own
+2 -2
View File
@@ -20,9 +20,9 @@ jobs:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: build
+7 -7
View File
@@ -20,9 +20,9 @@ jobs:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: Install goimports
@@ -42,7 +42,7 @@ jobs:
- name: golangci-lint
uses: golangci/golangci-lint-action@v9
with:
version: v2.12
version: v2.5
test:
name: Test ${{ matrix.name }}
@@ -80,9 +80,9 @@ jobs:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: Build
@@ -125,9 +125,9 @@ jobs:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
- uses: actions/setup-go@v6
with:
go-version: '1.26'
go-version: '1.25'
check-latest: true
- name: Build ${{ matrix.name }}
-82
View File
@@ -7,88 +7,6 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [Unreleased]
## [1.11.0] - 2026-07-23
See the [v1.11.0](https://github.com/slackhq/nebula/milestone/25?closed=1) milestone for a complete list of changes.
### Breaking
- Logging has switched from logrus to Go's structured `slog`. Log output changes: levels are upper case
(`level=INFO`), trace prints as `level=DEBUG-4`, timestamps are always RFC3339Nano and `logging.timestamp_format`
is ignored, and some messages were reworded. Review any log parsing before upgrading. This is also an API break
for embedders, as constructors now take a `*slog.Logger`. (#1672, #1734, #1621)
- `firewall.inbound_action` and `firewall.outbound_action` (used to set reject vs. drop policy) were each being
applied to the opposite direction, that is now corrected. This only affects how blocked packets are answered, not
which packets the firewall allows or denies. If you set either of these you are getting the behavior of the other
one today and likely want to swap them before upgrading. (#1798)
- On Windows, Nebula now installs WFP PERMIT filters for the nebula adapter and the listener port by default. WFP
sits below Windows Defender Firewall, so any WDF inbound rules you rely on for either will no longer apply. Set
`tun.windows_bypass_wdf` and `listen.windows_bypass_wdf` to false to leave WDF in charge. (#1710)
- On Windows, the nebula device is now set to the `private` network category instead of whatever Windows decided,
which is usually `Public`. This makes the host firewall less restrictive on the overlay. Set
`tun.network_category` to `unset` to keep the old behavior. (#1710)
- Reject packets for non-TCP now use ICMP code 13, communication administratively prohibited, instead of code 3,
port unreachable. Anything keying off the old code needs updating. (#1766, #1768)
- The SSH debug server's profiling commands are now confined to `sshd.sandbox_dir`, which defaults to
`$TMP/nebula-debug`. Relative paths resolve inside it and absolute paths outside it are rejected, so anything
scripting `start-cpu-profile`, `save-heap-profile`, or `save-mutex-profile` with a path elsewhere needs the
directory set. The directory is not created for you. (#1622)
### Added
- Sign the Windows release binaries. (#1718)
- Generate IPv6 reject packets, matching the existing IPv4 behavior. (#1766, #1767, #1768)
- Accept `-` in `nebula-cert` to read from stdin or write to stdout. (#1714)
- Search for both `config.yml` and `config.yaml` in service and command line modes. (#1717)
- Add version labels to the Docker/OCI images. (#1772)
- Rebind the listener and re-query lighthouses on macOS when the underlay network changes, so devices moving
between wifi and wired or between networks recover without waiting for dead tunnel detection. Controlled by
`listen.rebind_on_network_change` (default `true`, not reloadable). (#1816)
### Changed
- Reload the firewall when the unsafe networks in the certificate change. (#1719)
- Reconfigure, start, and stop the stats listener on a config reload instead of requiring a restart. (#1670)
- Update a static host's addresses when they change on reload. (#1713)
- Don't require a port on ICMP firewall rules. (#1609)
- Connection track ICMP traffic. (#1602)
- Return `NODATA` instead of `NXDOMAIN` from the DNS server for a name that exists but has no record of the
requested type, so clients that query `AAAA` first (busybox/Alpine) fall through to `A`. (#1668)
- Record the local host's details in the DNS server. (#1716)
- Install Windows unsafe routes as link routes. (#1709)
- Reduce relay handshake log spam, and only log a handshake send error at error level when the remote list
changes. (#1733, #1765, #1810)
- Start, stop, and reload subsystems (DNS, stats, conntrack, ssh, punchy) cleanly without leaking goroutines. (#1640, #1654, #1661, #1667, #1669, #1708, #1806, #1815)
- `Control` is now safe to stop and wait on from any lifecycle state, and a new `Control.Wait` blocks until nebula
has fully stopped and returns the first fatal reader error. Failed starts release the udp sockets and tun fd
instead of leaking them. (#1794)
- Trigger an immediate lighthouse update when reconnecting to or adding a lighthouse instead of waiting for the next update tick. (#1645)
- Bring the Darwin and OpenBSD tun implementations in line with the other BSDs. (#1703)
- Update to build against go v1.26. (#1818)
- Various dependency updates. (#1586, #1587, #1604, #1617, #1618, #1627, #1628, #1629, #1652, #1664, #1665, #1697, #1721, #1732, #1742, #1743, #1750, #1763, #1771, #1782, #1800, #1807)
### Fixed
- Fix a data race on a host's remote address that could send packets to the wrong address during a roam. (#1773)
- Fix tunnels that could permanently escape connection manager monitoring. (#1752)
- Fix a crash when reloading the SSH server's trusted keys. (#1787)
- Fix hostmap corruption when a host has multiple overlay addresses. Each address now gets its own list instead of
a single shared chain, which also fixes two latent bugs on the add and makePrimary paths. (#1788, #1790)
- Apply `remote_allow_list` IPv4 rules to 4-in-6 mapped addresses. (#1786)
- Don't panic in the DNS server on a short or empty query name. (#1635)
- Advance the replay window on relayed packets so a relay drops replayed frames instead of re-forwarding them. (#1751)
- Fix a race in relay state handling. (#1753)
- Lock replay window updates so concurrent readers can't corrupt it. (#1802)
- Reject malformed handshakes more reliably, including invalid ed25519 key lengths. (#1601, #1756)
- Properly handle `closetunnel` packets. (#1638)
- Fix an IPv6 extension-header length overflow that could make the firewall parse the wrong protocol and ports. (#1789)
- Fix relay re-establishment when a handshake arrives over a relay entry that a one-sided teardown left
`Disestablished`, which silently dropped every send until dead tunnel detection forced a re-handshake. (#1805)
- Don't build new relay state on a tunnel that was just discarded. (#1796)
- Don't delete the wrong pending hostinfo in the handshake manager. (#1811)
- Don't call the packet reader after a UDP error on Darwin. (#1755)
- Open the FreeBSD tun device non blocking. (#1666)
## [1.10.3] - 2026-02-06
### Security
+1 -5
View File
@@ -161,10 +161,6 @@ bin-pkcs11: BUILD_ARGS += -tags pkcs11
bin-pkcs11: CGO_ENABLED = 1
bin-pkcs11: bin
# Build with the pprof debug server (serves on :6060). See startPprofServer.
debug: BUILD_ARGS += -tags debug
debug: bin
bin:
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
@@ -284,5 +280,5 @@ smoke-vagrant/%: bin-docker build/%/nebula
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
.FORCE:
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin debug build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
.DEFAULT_GOAL := bin
+4 -8
View File
@@ -53,12 +53,7 @@ func main() {
l := logging.NewLogger(os.Stdout)
if *serviceFlag != "" {
if *configTest {
fmt.Println("-test is not supported with -service, run the config test without -service")
os.Exit(1)
}
if err := doService(configPath, Build, serviceFlag); err != nil {
if err := doService(configPath, configTest, Build, serviceFlag); err != nil {
l.Error("Service command failed", "error", err)
os.Exit(1)
}
@@ -98,14 +93,15 @@ func main() {
}
if !*configTest {
if err := ctrl.Start(); err != nil {
wait, err := ctrl.Start()
if err != nil {
util.LogWithContextIfNeeded("Error while running", err, l)
os.Exit(1)
}
go ctrl.ShutdownBlock()
if err := ctrl.Wait(); err != nil {
if err := wait(); err != nil {
l.Error("Nebula stopped due to fatal error", "error", err)
os.Exit(2)
}
+6 -25
View File
@@ -3,7 +3,6 @@ package main
import (
"fmt"
"log"
"os"
"github.com/kardianos/service"
"github.com/slackhq/nebula"
@@ -15,6 +14,7 @@ var logger service.Logger
type program struct {
configPath *string
configTest *bool
build string
control *nebula.Control
}
@@ -40,41 +40,22 @@ func (p *program) Start(s service.Service) error {
}
})
p.control, err = nebula.Main(c, false, Build, l, nil)
p.control, err = nebula.Main(c, *p.configTest, Build, l, nil)
if err != nil {
return err
}
if err := p.control.Start(); err != nil {
return err
}
// Nebula can stop itself on a fatal packet reader error, make sure to log it if it happens.
go func() {
if err := p.control.Wait(); err != nil {
logger.Error(fmt.Sprintf("Nebula stopped due to fatal error: %v", err))
os.Exit(2)
}
}()
p.control.Start()
return nil
}
func (p *program) Stop(s service.Service) error {
logger.Info("Nebula service stopping.")
if p.control == nil {
return nil
}
p.control.Stop()
// block until nebula has fully drained before reporting stopped.
// error logging is handled by Start.
_ = p.control.Wait()
return nil
}
func doService(configPath *string, build string, serviceFlag *string) error {
func doService(configPath *string, configTest *bool, build string, serviceFlag *string) error {
if *configPath == "" {
p, err := config.DefaultPath()
if err != nil {
@@ -92,6 +73,7 @@ func doService(configPath *string, build string, serviceFlag *string) error {
prg := &program{
configPath: configPath,
configTest: configTest,
build: build,
}
@@ -123,9 +105,8 @@ func doService(configPath *string, build string, serviceFlag *string) error {
switch *serviceFlag {
case "run":
if err := s.Run(); err != nil {
// Route any errors to the system logger and report the failure
// Route any errors to the system logger
logger.Error(err)
return err
}
default:
if err := service.Control(s, *serviceFlag); err != nil {
-96
View File
@@ -1,96 +0,0 @@
//go:build linux && !android && !e2e_testing
package main
import (
"fmt"
"net/netip"
"os"
"path/filepath"
"runtime"
"testing"
"time"
"github.com/slackhq/nebula"
"github.com/slackhq/nebula/cert"
cert_test "github.com/slackhq/nebula/cert_test"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/test"
"github.com/stretchr/testify/require"
)
// TestControlStopClosesOnTimer reproduces the dnclient lifecycle: nebula runs as
// a library, and on a config update dnclient calls Stop() in-process to tear the
// old instance down before starting a new one. This boots a real nebula (real
// blocking UDP sockets, tun disabled), lets it run, then Stop()s it on a timer
// and asserts it actually closes. If the reader goroutines parked in recvmmsg
// don't wake on Close(), Wait() blocks forever and this fails with a goroutine
// dump instead of relying on a process signal to unstick them.
func TestControlStopClosesOnTimer(t *testing.T) {
l := test.NewLogger()
dir := t.TempDir()
before := time.Now().Add(-time.Hour)
after := time.Now().Add(time.Hour)
ca, _, caKey, caPEM := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, before, after, nil, nil, nil)
networks := []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")}
_, _, keyPEM, certPEM := cert_test.NewTestCert(cert.Version2, cert.Curve_CURVE25519, ca, caKey, "close-on-timer", before, after, networks, nil, nil)
caPath := filepath.Join(dir, "ca.pem")
certPath := filepath.Join(dir, "cert.pem")
keyPath := filepath.Join(dir, "key.pem")
require.NoError(t, os.WriteFile(caPath, caPEM, 0o600))
require.NoError(t, os.WriteFile(certPath, certPEM, 0o600))
require.NoError(t, os.WriteFile(keyPath, keyPEM, 0o600))
// tun disabled so no device/root is needed; routines: 2 so we exercise the
// multi-socket (SO_REUSEPORT) teardown, which is where dnclient runs.
configBody := fmt.Sprintf(`
pki:
ca: %s
cert: %s
key: %s
listen:
host: 127.0.0.1
port: 0
tun:
disabled: true
firewall:
outbound:
- port: any
proto: any
host: any
inbound:
- port: any
proto: any
host: any
routines: 2
`, caPath, certPath, keyPath)
require.NoError(t, os.WriteFile(filepath.Join(dir, "config.yml"), []byte(configBody), 0o600))
c := config.NewC(l)
require.NoError(t, c.Load(dir))
ctrl, err := nebula.Main(c, false, "close-on-timer", l, nil)
require.NoError(t, err)
require.NoError(t, ctrl.Start())
// Run like a live nebula, then close on a timer, exactly as dnclient does.
<-time.NewTimer(5 * time.Second).C
stopped := make(chan struct{})
go func() {
ctrl.Stop() // closes the udp sockets (shutdown(2)) and the tun
ctrl.Wait() // blocks until every reader goroutine has returned
close(stopped)
}()
select {
case <-stopped:
t.Log("nebula closed cleanly on timer")
case <-time.After(10 * time.Second):
buf := make([]byte, 1<<20)
n := runtime.Stack(buf, true)
t.Fatalf("nebula did NOT close within 10s of Stop(): a blocking reader never woke\n%s", buf[:n])
}
}
+3 -2
View File
@@ -84,7 +84,8 @@ func main() {
}
if !*configTest {
if err := ctrl.Start(); err != nil {
wait, err := ctrl.Start()
if err != nil {
util.LogWithContextIfNeeded("Error while running", err, l)
os.Exit(1)
}
@@ -92,7 +93,7 @@ func main() {
go ctrl.ShutdownBlock()
notifyReady(l)
if err := ctrl.Wait(); err != nil {
if err := wait(); err != nil {
l.Error("Nebula stopped due to fatal error", "error", err)
os.Exit(2)
}
-1
View File
@@ -25,7 +25,6 @@ func newTestLighthouse() *LightHouse {
lighthouses := []netip.Addr{}
staticList := map[netip.Addr]struct{}{}
lh.localAddrsFn = func(*LocalAllowList) []netip.Addr { return nil }
lh.lighthouses.Store(&lighthouses)
lh.staticList.Store(&staticList)
+1 -66
View File
@@ -2,25 +2,15 @@ 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
const ReplayWindow = 1024
type ConnectionState struct {
eKey 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
}
+24 -58
View File
@@ -53,7 +53,6 @@ type Control struct {
statsStart func()
dnsStart func()
lighthouseStart func()
networkChangeStart func(rebind func())
connectionManagerStart func(context.Context)
}
@@ -70,29 +69,29 @@ type ControlHostInfo struct {
}
// Start actually runs nebula, this is a nonblocking call.
// Use Wait to block until nebula has fully stopped and to learn whether a fatal reader error caused the shutdown.
func (c *Control) Start() error {
// The returned function blocks until nebula has fully stopped and returns the
// first fatal reader error (if any). A nil error means nebula shut down
// gracefully; a non-nil error means a reader hit an unexpected failure that
// triggered the shutdown.
func (c *Control) Start() (func() error, error) {
c.stateLock.Lock()
defer c.stateLock.Unlock()
switch c.state {
case StateReady:
//yay!
case StateStopped, StateStopping:
return ErrAlreadyStopped
return nil, ErrAlreadyStopped
case StateStarted:
return ErrAlreadyStarted
return nil, ErrAlreadyStarted
default:
return ErrUnknownState
return nil, ErrUnknownState
}
// Activate the interface
err := c.f.activate()
if err != nil {
// Cancel before Close so a caller returning from Wait always observes a dead Context
c.cancel()
_ = c.f.Close()
c.state = StateStopped
return err
return nil, err
}
// Call all the delayed funcs that waited patiently for the interface to be created.
@@ -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)
}
@@ -118,9 +114,13 @@ func (c *Control) Start() error {
c.f.triggerShutdown = c.Stop
// Start reading packets.
c.f.run()
out, err := c.f.run()
if err != nil {
c.state = StateStopped
return nil, err
}
c.state = StateStarted
return nil
return out, nil
}
func (c *Control) State() RunState {
@@ -133,26 +133,10 @@ func (c *Control) Context() context.Context {
return c.ctx
}
// Stop tears nebula down, closing all tunnels and releasing everything it holds.
// Use Wait to block until the shutdown has completed.
// A Control that has been stopped cannot be started again, Start will return ErrAlreadyStopped.
// Stop is a non-blocking call that signals nebula to close all tunnels and shut down
func (c *Control) Stop() {
c.stateLock.Lock()
switch c.state {
case StateStarted:
// Fall through to the full teardown below
case StateReady:
// Never started
c.cancel()
c.state = StateStopped
if err := c.f.Close(); err != nil {
c.l.Error("Close interface failed", "error", err)
}
c.stateLock.Unlock()
return
default:
if c.state != StateStarted {
c.stateLock.Unlock()
// We are stopping or stopped already
return
@@ -161,26 +145,19 @@ func (c *Control) Stop() {
c.state = StateStopping
c.stateLock.Unlock()
// Closing tunnels can be slow with a large hostmap, don't hold the lock for it
// Stop the handshakeManager (and other services), to prevent new tunnels from
// being created while we're shutting them all down.
c.cancel()
c.CloseAllTunnels(false)
c.stateLock.Lock()
c.state = StateStopped
c.CloseAllTunnels(false)
if err := c.f.Close(); err != nil {
c.l.Error("Close interface failed", "error", err)
}
c.stateLock.Lock()
c.state = StateStopped
c.stateLock.Unlock()
}
// Wait blocks until nebula has fully stopped, either via Stop or an internal fatal error,
// and returns the first fatal packet reader error if there was one.
// It is safe to call from multiple goroutines and at any point in the lifecycle,
// but a Wait on a Control that is never started and never stopped will block forever.
func (c *Control) Wait() error {
return c.f.wait()
}
// ShutdownBlock will listen for and block on term and interrupt signals, calling Control.Stop() once signalled
func (c *Control) ShutdownBlock() {
sigChan := make(chan os.Signal, 1)
@@ -193,20 +170,9 @@ func (c *Control) ShutdownBlock() {
c.Stop()
}
// RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change.
// RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change
func (c *Control) RebindUDPServer() {
c.stateLock.Lock()
defer c.stateLock.Unlock()
if c.state != StateStarted {
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()
-309
View File
@@ -1,309 +0,0 @@
package nebula
import (
"context"
"errors"
"io"
"net/netip"
"sync"
"testing"
"time"
"github.com/gaissmai/bart"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/overlay/batch"
"github.com/slackhq/nebula/overlay/tio"
"github.com/slackhq/nebula/routing"
"github.com/slackhq/nebula/test"
"github.com/slackhq/nebula/udp"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
type fakeDevice struct {
closeOnce sync.Once
closedCh chan struct{}
closed bool
}
func newFakeDevice() *fakeDevice {
return &fakeDevice{closedCh: make(chan struct{})}
}
// Read blocks until Close like a real tun with no traffic, then reports EOF
// the same way a closed device does
func (d *fakeDevice) Read() ([]tio.Packet, error) {
<-d.closedCh
return nil, io.EOF
}
func (d *fakeDevice) Write(p []byte) (int, error) { return len(p), nil }
func (d *fakeDevice) Close() error {
d.closeOnce.Do(func() {
d.closed = true
close(d.closedCh)
})
return nil
}
func (d *fakeDevice) Activate() error { return nil }
func (d *fakeDevice) Networks() []netip.Prefix { return nil }
func (d *fakeDevice) Name() string { return "fake" }
func (d *fakeDevice) RoutesFor(netip.Addr) routing.Gateways { return nil }
func (d *fakeDevice) Queues(int) ([]tio.Queue, error) { return []tio.Queue{d}, nil }
// newReadyControl hand-builds the minimum Control that Main would have
// produced right before Start, including the construction token NewInterface
// takes so waiters block until Close releases the resources
func newReadyControl(t *testing.T) (*Control, *fakeDevice, *fakeConn) {
l := test.NewLogger()
dev := newFakeDevice()
conn := &fakeConn{}
ctx, cancel := context.WithCancel(context.Background())
myVpnNet := netip.MustParsePrefix("10.128.0.1/16")
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
lh, err := NewLightHouseFromConfig(ctx, l, config.NewC(l), cs, nil, nil)
require.NoError(t, err)
f := &Interface{
ctx: ctx,
inside: dev,
outside: conn,
writers: []udp.Conn{conn},
batchers: make([]*batch.MultiCoalescer, 1),
routines: 1,
hostMap: newHostMap(l),
lightHouse: lh,
l: l,
}
f.wg.Add(1)
return &Control{
state: StateReady,
f: f,
l: l,
ctx: ctx,
cancel: cancel,
}, dev, conn
}
func TestControl_StopBeforeStart(t *testing.T) {
c, dev, conn := newReadyControl(t)
// A Stop on a never started control must release everything Main acquired
c.Stop()
assert.Equal(t, StateStopped, c.State())
assert.True(t, dev.closed, "the tun device should have been closed")
assert.True(t, conn.closed, "the udp socket should have been closed")
require.ErrorIs(t, c.ctx.Err(), context.Canceled, "the service context should have been cancelled")
// Wait must return promptly now that the resources are released
require.NoError(t, c.Wait())
// A stopped control can never be started
err := c.Start()
require.ErrorIs(t, err, ErrAlreadyStopped)
// A second Stop is a harmless no-op
c.Stop()
assert.Equal(t, StateStopped, c.State())
require.NoError(t, c.Wait())
}
func TestControl_WaitBlocksUntilStop(t *testing.T) {
c, _, _ := newReadyControl(t)
done := make(chan error, 1)
go func() { done <- c.Wait() }()
select {
case <-done:
t.Fatal("Wait returned before Stop")
case <-time.After(50 * time.Millisecond):
}
c.Stop()
select {
case err := <-done:
require.NoError(t, err)
case <-time.After(time.Second):
t.Fatal("Wait did not return after Stop")
}
}
type fakeConn struct {
closed bool
rebinds int
}
func (c *fakeConn) Rebind() error { c.rebinds++; return nil }
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
func (c *fakeConn) ListenOut(_ udp.EncReader, _ func()) error { return nil }
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
func (c *fakeConn) WriteBatch(bufs [][]byte, _ []netip.AddrPort) (int, error) {
return len(bufs), nil
}
func (c *fakeConn) ReloadConfig(_ *config.C) {}
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
func (c *fakeConn) Close() error { c.closed = true; return nil }
type multiqueueDevice struct {
*fakeDevice
}
// Queues claims multiqueue support but fails to open the second queue,
// exercising the activation error path.
func (d *multiqueueDevice) Queues(n int) ([]tio.Queue, error) {
if n > 1 {
return nil, errors.New("second queue failed to open")
}
return d.fakeDevice.Queues(n)
}
func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
dev := &multiqueueDevice{fakeDevice: newFakeDevice()}
conn := &fakeConn{}
ctx, cancel := context.WithCancel(context.Background())
f := &Interface{
ctx: ctx,
inside: dev,
outside: conn,
writers: []udp.Conn{conn},
batchers: make([]*batch.MultiCoalescer, 2),
routines: 2,
l: test.NewLogger(),
}
f.wg.Add(1)
c := &Control{
state: StateReady,
f: f,
l: test.NewLogger(),
ctx: ctx,
cancel: cancel,
}
// The second reader fails to open, everything must be released
err := c.Start()
require.Error(t, err)
assert.Equal(t, StateStopped, c.State())
assert.True(t, dev.closed, "the tun device should have been closed")
assert.True(t, conn.closed, "the udp socket should have been closed")
require.ErrorIs(t, c.ctx.Err(), context.Canceled)
// And Wait must not hang on the construction token
require.NoError(t, c.Wait())
}
func TestInterface_CloseIsIdempotent(t *testing.T) {
dev := newFakeDevice()
f := &Interface{
inside: dev,
l: test.NewLogger(),
}
f.wg.Add(1)
require.NoError(t, f.Close())
assert.True(t, dev.closed)
// A second Close must not double release the wg token or the device
require.NoError(t, f.Close())
require.NoError(t, f.wait())
}
func TestControl_FatalErrorReportsThroughWait(t *testing.T) {
c, dev, conn := newReadyControl(t)
// Mirror what Start wires up, without needing real packet readers
c.f.triggerShutdown = c.Stop
c.state = StateStarted
boom := errors.New("boom")
c.f.onFatal(boom)
require.ErrorIs(t, c.Wait(), boom)
assert.Equal(t, StateStopped, c.State())
assert.True(t, dev.closed)
assert.True(t, conn.closed)
// A second fatal error must not fire the shutdown again or replace the first
c.f.onFatal(errors.New("later"))
require.ErrorIs(t, c.Wait(), boom)
// Wait stays factual, a Stop after the death does not mask the error
c.Stop()
require.ErrorIs(t, c.Wait(), boom)
}
func TestControl_ConcurrentStopAndStart(t *testing.T) {
c, _, _ := newReadyControl(t)
var wg sync.WaitGroup
for i := 0; i < 2; i++ {
wg.Go(func() { c.Stop() })
}
wg.Go(func() { _ = c.Start() })
wg.Go(func() {
_ = c.Wait()
// A returned Wait must always observe the final state, no matter how
// the race resolved
assert.Equal(t, StateStopped, c.State())
})
wg.Wait()
// However the race resolves, the control must end fully stopped with no
// panic and Wait must observe the final state
require.NoError(t, c.Wait())
assert.Equal(t, StateStopped, c.State())
err := c.Start()
require.ErrorIs(t, err, ErrAlreadyStopped)
}
func TestControl_StartStopLifecycle(t *testing.T) {
c, dev, conn := newReadyControl(t)
err := c.Start()
require.NoError(t, err)
assert.Equal(t, StateStarted, c.State())
err = c.Start()
require.ErrorIs(t, err, ErrAlreadyStarted)
// Stop must unpark the reader blocked in the device and release everything
c.Stop()
assert.Equal(t, StateStopped, c.State())
assert.True(t, dev.closed, "the tun device should have been closed")
assert.True(t, conn.closed, "the udp socket should have been closed")
require.ErrorIs(t, c.ctx.Err(), context.Canceled)
// The reader drained off a closed device, that is not a fatal error
require.NoError(t, c.Wait())
err = c.Start()
require.ErrorIs(t, err, ErrAlreadyStopped)
}
func TestControl_RebindIsGatedByState(t *testing.T) {
c, _, conn := newReadyControl(t)
// A rebind before Start reaches nothing, the interface is not up
c.RebindUDPServer()
assert.Equal(t, 0, conn.rebinds, "rebind before start must be a no-op")
err := c.Start()
require.NoError(t, err)
c.RebindUDPServer()
assert.Equal(t, 1, conn.rebinds, "rebind while started must reach the conn")
// A rebind racing a completed stop must not touch the closed conn
c.Stop()
require.NoError(t, c.Wait())
c.RebindUDPServer()
assert.Equal(t, 1, conn.rebinds, "rebind after stop must be a no-op")
}
+1 -21
View File
@@ -108,19 +108,7 @@ func (c *Control) GetVpnAddrs() []netip.Addr {
}
func (c *Control) GetUDPAddr() netip.AddrPort {
return c.f.outside.(*udp.TesterConn).GetAddr()
}
// SetUDPAddr moves this node to a new underlay address, standing in for a laptop waking up on a different
// network. Register the new address with the router as well or nothing will route back.
func (c *Control) SetUDPAddr(addr netip.AddrPort) {
c.f.outside.(*udp.TesterConn).SetAddr(addr)
}
// SetLocalAddrsFn replaces underlay address discovery so a test can advertise its simulated address instead of
// whatever this machine's NICs happen to be. Call it before Start, SendUpdate reads it from the update worker.
func (c *Control) SetLocalAddrsFn(fn func(*LocalAllowList) []netip.Addr) {
c.f.lightHouse.localAddrsFn = fn
return c.f.outside.(*udp.TesterConn).Addr
}
func (c *Control) KillPendingTunnel(vpnIp netip.Addr) bool {
@@ -137,14 +125,6 @@ func (c *Control) GetHostmap() *HostMap {
return c.f.hostMap
}
// GetHostmapIndexCount returns the number of entries in the main hostmap Indexes table, holding
// the hostmap read lock so tests can poll it while connection manager churns tunnels.
func (c *Control) GetHostmapIndexCount() int {
c.f.hostMap.RLock()
defer c.f.hostMap.RUnlock()
return len(c.f.hostMap.Indexes)
}
func (c *Control) GetF() *Interface {
return c.f
}
-187
View File
@@ -1,187 +0,0 @@
// Package cpupick chooses which CPUs the tun reader threads pin to when the
// operator has not chosen for us (tun.cpu_affinity). The stock spread —
// allowed[i] for routine i — has two failure modes this package exists to fix:
//
// - every co-located nebula starts its spread at allowed[0], so N instances
// on one box stack their readers onto the same cores, and allowed[0] is
// usually CPU 0, the core housekeeping and default IRQ affinity already
// favor;
// - on heterogeneous CPUs (ARM big.LITTLE, Intel P/E hybrids, AMD compact
// cores) low IDs are not necessarily fast cores, and pinning an encrypt
// thread to an efficiency core caps that queue's throughput.
//
// Default instead returns a preference-ordered pin list: the allowed set
// filtered to performance cores (when the platform distinguishes them and
// enough remain for every routine), confined to a single NUMA node and spread
// across distinct physical cores when the topology permits, CPU 0's physical
// core demoted to last resort, and the order rotated by a stable per-instance
// key so co-located instances spread instead of stacking.
package cpupick
import (
"log/slog"
"github.com/slackhq/nebula/util"
)
// topology is the slice of machine layout arrange consults: the NUMA node
// and the physical core behind each candidate CPU, plus which core CPU 0
// lives on (zeroCore, -1 when unknown — tracked separately because CPU 0's
// SMT sibling deserves demotion even when CPU 0 itself isn't a candidate).
// Probed from sysfs on Linux; flatTopology stands in when the platform can't
// say, which turns every topology rule into a no-op rather than a wrong
// answer.
type topology struct {
nodeOf map[int]int
coreOf map[int]int
zeroCore int
}
// flatTopology places every CPU on node 0 and on a physical core of its own.
func flatTopology(cpus []int) topology {
t := topology{
nodeOf: make(map[int]int, len(cpus)),
coreOf: make(map[int]int, len(cpus)),
zeroCore: -1,
}
for i, c := range cpus {
t.nodeOf[c] = 0
t.coreOf[c] = i
if c == 0 {
t.zeroCore = i
}
}
return t
}
// Default computes the pin order for `routines` tun readers. key is any
// stable per-instance value; the bound UDP port is ideal — distinct across
// co-located instances, stable across restarts so benchmark runs stay
// comparable. Returns nil when there is nothing useful to say (no affinity
// support on this platform, lookup failure); callers keep their existing
// fallback spread.
func Default(routines int, key uint64, l *slog.Logger) []int {
allowed, err := util.AllowedCPUs()
if err != nil || len(allowed) == 0 {
return nil
}
perf, signal := perfCPUs(allowed)
cands := pickCandidates(allowed, perf, routines)
if len(cands) == 0 {
return nil
}
if len(perf) < routines {
signal = ""
}
cpus := arrange(cands, readTopology(cands), routines, splitmix64(key))
if l != nil {
l.Info("chose default pin CPUs for tun readers",
"cpus", cpus[:min(routines, len(cpus))],
"perfSignal", signal)
}
return cpus
}
// pickCandidates applies the enough-for-everyone guard: a perf filter that
// leaves fewer candidates than routines is discarded — giving every reader
// its own (possibly slow) core beats stacking two readers on a fast one.
func pickCandidates(allowed, perf []int, routines int) []int {
if len(perf) < routines {
return allowed
}
return perf
}
// arrange turns the candidate set into the final pin order:
//
// 1. NUMA: when at least one node holds enough candidates for every
// routine, confine to one such node, chosen by the instance hash. The
// readers share hostmap and cipher state, so splitting one instance
// across nodes taxes every packet — and co-located instances that hash
// to different nodes stop competing entirely. When no node is big
// enough, span nodes rather than stack readers.
// 2. Rotate the preferred candidates by the hash so instances spread.
// 3. SMT: emit one thread per physical core before any of their siblings —
// two encrypt threads on one core split its execution units. Siblings
// still follow for the routines > cores case.
// 4. CPU 0's whole physical core goes last: housekeeping and default IRQ
// noise on CPU 0 bleeds into its SMT sibling too. Within that tail the
// sibling precedes CPU 0 itself, which only catches the bleed-through.
//
// The rotation happens before the SMT pass so each instance's one-per-core
// walk also starts at a different core, and CPU 0's core is excluded from
// the rotation so no hash value can put it back at the front.
func arrange(cands []int, topo topology, routines int, h uint64) []int {
byNode := map[int][]int{}
var nodes []int
for _, c := range cands {
n := topo.nodeOf[c]
if _, ok := byNode[n]; !ok {
nodes = append(nodes, n)
}
byNode[n] = append(byNode[n], c)
}
var eligible []int
for _, n := range nodes {
if len(byNode[n]) >= routines {
eligible = append(eligible, n)
}
}
if len(eligible) > 0 {
cands = byNode[eligible[int(h%uint64(len(eligible)))]]
}
// Split off CPU 0's core: its siblings tail the list, CPU 0 tails them.
preferred := make([]int, 0, len(cands))
var zeroTail []int
hasZero := false
for _, c := range cands {
switch {
case c == 0:
hasZero = true
case topo.zeroCore >= 0 && topo.coreOf[c] == topo.zeroCore:
zeroTail = append(zeroTail, c)
default:
preferred = append(preferred, c)
}
}
if hasZero {
zeroTail = append(zeroTail, 0)
}
if len(preferred) == 0 {
return zeroTail // CPU 0's core is all we have
}
// The node pick consumed the low hash bits; rotate by the high ones so
// the two choices stay independent.
off := int((h >> 32) % uint64(len(preferred)))
rot := make([]int, 0, len(preferred))
rot = append(rot, preferred[off:]...)
rot = append(rot, preferred[:off]...)
seenCore := make(map[int]bool, len(rot))
out := make([]int, 0, len(cands))
var siblings []int
for _, c := range rot {
g := topo.coreOf[c]
if seenCore[g] {
siblings = append(siblings, c)
continue
}
seenCore[g] = true
out = append(out, c)
}
out = append(out, siblings...)
out = append(out, zeroTail...)
return out
}
// splitmix64 decorrelates instance keys before the selection modulos: ports
// on one box often share spacing (4242/4243, or round steps like +1000) that
// raw key%len arithmetic would fold onto the same offset.
func splitmix64(x uint64) uint64 {
x += 0x9e3779b97f4a7c15
x = (x ^ (x >> 30)) * 0xbf58476d1ce4e5b9
x = (x ^ (x >> 27)) * 0x94d049bb133111eb
return x ^ (x >> 31)
}
-171
View File
@@ -1,171 +0,0 @@
package cpupick
import (
"slices"
"testing"
)
// pairTopo builds a topology where consecutive candidate pairs are SMT
// siblings: (cpus[0],cpus[1]) share a core, (cpus[2],cpus[3]) the next, ...
// All CPUs land on node 0.
func pairTopo(cpus []int) topology {
t := topology{
nodeOf: make(map[int]int, len(cpus)),
coreOf: make(map[int]int, len(cpus)),
zeroCore: -1,
}
for i, c := range cpus {
t.nodeOf[c] = 0
t.coreOf[c] = i / 2
if c == 0 {
t.zeroCore = i / 2
}
}
return t
}
func TestArrangeDemotesZeroForEveryKey(t *testing.T) {
candidates := []int{0, 1, 2, 3, 4, 5, 6, 7}
for key := range uint64(64) {
got := arrange(candidates, flatTopology(candidates), 4, splitmix64(key))
if len(got) != len(candidates) {
t.Fatalf("key %d: len=%d want %d", key, len(got), len(candidates))
}
if got[0] == 0 {
t.Errorf("key %d: CPU 0 at the front: %v", key, got)
}
if got[len(got)-1] != 0 {
t.Errorf("key %d: CPU 0 not demoted to last: %v", key, got)
}
sorted := slices.Clone(got)
slices.Sort(sorted)
if !slices.Equal(sorted, candidates) {
t.Errorf("key %d: not a permutation: %v", key, got)
}
}
}
func TestArrangeDemotesZeroSiblings(t *testing.T) {
// Pairs (0,1),(2,3),(4,5),(6,7): CPU 0's core — 0 and its sibling 1 —
// must tail the list, sibling ahead of 0 itself.
candidates := []int{0, 1, 2, 3, 4, 5, 6, 7}
for key := range uint64(64) {
got := arrange(candidates, pairTopo(candidates), 2, splitmix64(key))
n := len(got)
if got[n-1] != 0 || got[n-2] != 1 {
t.Fatalf("key %d: tail = %v, want [... 1 0]", key, got)
}
}
}
func TestArrangeZeroSiblingWithoutZero(t *testing.T) {
// CPU 0 excluded (cpuset) but its sibling 1 remains: the sibling still
// tails the list when the topology knows which core CPU 0 lives on.
candidates := []int{1, 2, 3, 4, 5}
topo := pairTopo([]int{0, 1, 2, 3, 4, 5})
got := arrange(candidates, topo, 2, splitmix64(7))
if got[len(got)-1] != 1 {
t.Errorf("CPU 0's sibling not demoted: %v", got)
}
}
func TestArrangeRotatesByKey(t *testing.T) {
candidates := []int{1, 2, 3, 4, 5, 6, 7, 8}
seen := map[int]bool{}
for key := range uint64(64) {
seen[arrange(candidates, flatTopology(candidates), 4, splitmix64(key))[0]] = true
}
// 64 hashed keys over 8 slots must hit more than one starting CPU, or
// co-located instances would all stack again.
if len(seen) < 2 {
t.Errorf("rotation never varied across keys: %v", seen)
}
}
func TestArrangeStableForSameKey(t *testing.T) {
candidates := []int{0, 2, 4, 6}
topo := flatTopology(candidates)
a := arrange(candidates, topo, 2, splitmix64(4242))
b := arrange(candidates, topo, 2, splitmix64(4242))
if !slices.Equal(a, b) {
t.Errorf("same key ordered differently: %v vs %v", a, b)
}
}
func TestArrangeZeroOnly(t *testing.T) {
if got := arrange([]int{0}, flatTopology([]int{0}), 1, splitmix64(7)); !slices.Equal(got, []int{0}) {
t.Errorf("sole CPU 0 must survive: %v", got)
}
}
func TestArrangeSMTSiblingsLast(t *testing.T) {
// Pairs (1,2),(3,4),(5,6),(7,8): the first four picks must cover four
// distinct physical cores before any sibling repeats.
candidates := []int{1, 2, 3, 4, 5, 6, 7, 8}
topo := pairTopo(candidates)
for key := range uint64(16) {
got := arrange(candidates, topo, 4, splitmix64(key))
seen := map[int]bool{}
for _, c := range got[:4] {
g := topo.coreOf[c]
if seen[g] {
t.Fatalf("key %d: sibling before all cores covered: %v", key, got)
}
seen[g] = true
}
}
}
func TestArrangeNUMAConfinesToOneNode(t *testing.T) {
// Two nodes of four; both fit routines=3, so the result must sit
// entirely inside one of them, and the hash must pick both across keys.
candidates := []int{1, 2, 3, 4, 10, 11, 12, 13}
topo := flatTopology(candidates)
for _, c := range []int{10, 11, 12, 13} {
topo.nodeOf[c] = 1
}
nodesSeen := map[int]bool{}
for key := range uint64(32) {
got := arrange(candidates, topo, 3, splitmix64(key))
if len(got) != 4 {
t.Fatalf("key %d: not confined to one node: %v", key, got)
}
n := topo.nodeOf[got[0]]
for _, c := range got {
if topo.nodeOf[c] != n {
t.Fatalf("key %d: spans nodes: %v", key, got)
}
}
nodesSeen[n] = true
}
if len(nodesSeen) != 2 {
t.Errorf("hash never spread instances across nodes: %v", nodesSeen)
}
}
func TestArrangeNUMASpansWhenNoNodeFits(t *testing.T) {
candidates := []int{1, 2, 3, 4, 10, 11, 12, 13}
topo := flatTopology(candidates)
for _, c := range []int{10, 11, 12, 13} {
topo.nodeOf[c] = 1
}
got := arrange(candidates, topo, 6, splitmix64(1))
if len(got) != len(candidates) {
t.Errorf("undersized nodes must span, got %v", got)
}
}
func TestPickCandidates(t *testing.T) {
allowed := []int{0, 1, 2, 3, 4, 5, 6, 7}
perf := []int{4, 5}
// Enough perf cores for every routine: only they are used.
if got := pickCandidates(allowed, perf, 2); !slices.Equal(got, perf) {
t.Errorf("perf filter not applied: %v", got)
}
// Perf filter too small for the routine count: discarded, everyone
// gets their own core from the full allowed set.
if got := pickCandidates(allowed, perf, 4); !slices.Equal(got, allowed) {
t.Errorf("undersized perf filter not discarded: %v", got)
}
}
-154
View File
@@ -1,154 +0,0 @@
//go:build linux
package cpupick
import (
"fmt"
"os"
"path/filepath"
"strconv"
"strings"
)
// capacityKeepPct is the cpu_capacity admission threshold, relative to the
// fastest allowed core. LITTLE cores are normalized to ~250-400 of the big
// core's 1024 while mid cores sit at ~75%+, so half of max separates little
// from the rest without splitting prime from mid on three-tier parts.
const capacityKeepPct = 50
// freqKeepPct is the cpuinfo_max_freq admission threshold. Favored-core
// turbo skew is 2-4% and ARM mid-vs-prime ~12%, while E-cores, LITTLE
// cores, and AMD compact cores all sit >= 20% below their siblings' max.
const freqKeepPct = 85
// perfCPUs partitions allowed into the subset that are "performance" cores,
// consulting (in order of authority):
//
// 1. cpu_capacity — arch_topology's normalized per-CPU capacity, exposed on
// arm/arm64/riscv; the scheduler's own view of big vs LITTLE.
// 2. /sys/devices/cpu_core/cpus — the Intel hybrid P-core PMU mask, present
// only on P/E parts (x86 has no cpu_capacity) and naming P cores outright.
// 3. cpuinfo_max_freq — the cross-vendor fallback; catches AMD compact
// cores, which neither of the above covers.
//
// Returns allowed unchanged (signal "") when nothing distinguishes the
// cores: homogeneous parts, VMs without cpufreq, sysfs unavailable.
func perfCPUs(allowed []int) ([]int, string) {
return perfCPUsFrom("/sys/devices/system/cpu", "/sys/devices/cpu_core/cpus", allowed)
}
func perfCPUsFrom(cpuDir, intelCoreMask string, allowed []int) ([]int, string) {
if cpus, ok := byPerCPUValue(cpuDir, "cpu_capacity", allowed, capacityKeepPct); ok {
return cpus, "cpu_capacity"
}
if cpus, ok := byIntelCoreMask(intelCoreMask, allowed); ok {
return cpus, "intel_core_pmu"
}
if cpus, ok := byPerCPUValue(cpuDir, "cpufreq/cpuinfo_max_freq", allowed, freqKeepPct); ok {
return cpus, "max_freq"
}
return allowed, ""
}
// byPerCPUValue keeps the allowed CPUs whose per-CPU sysfs value is at least
// keepPct percent of the maximum across allowed. Inconclusive (ok=false)
// when any CPU is missing the file or when every value is equal.
func byPerCPUValue(cpuDir, file string, allowed []int, keepPct int) ([]int, bool) {
vals := make([]int, len(allowed))
minV, maxV := 0, 0
for i, cpu := range allowed {
v, err := readIntFile(filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), file))
if err != nil {
return nil, false
}
vals[i] = v
if i == 0 || v < minV {
minV = v
}
if v > maxV {
maxV = v
}
}
if minV == maxV {
return nil, false // homogeneous by this signal; try the next one
}
keep := make([]int, 0, len(allowed))
for i, cpu := range allowed {
if vals[i]*100 >= maxV*keepPct {
keep = append(keep, cpu)
}
}
return keep, true
}
// byIntelCoreMask keeps the allowed CPUs named by the hybrid P-core PMU
// mask. Inconclusive when the file is absent (non-hybrid x86, other arches)
// or no allowed CPU is in the mask (the process was deliberately confined
// to E-cores; nothing useful to prefer within that).
func byIntelCoreMask(maskPath string, allowed []int) ([]int, bool) {
b, err := os.ReadFile(maskPath)
if err != nil {
return nil, false
}
set, err := parseCPUList(strings.TrimSpace(string(b)))
if err != nil || len(set) == 0 {
return nil, false
}
pcore := make(map[int]bool, len(set))
for _, c := range set {
pcore[c] = true
}
keep := make([]int, 0, len(allowed))
for _, cpu := range allowed {
if pcore[cpu] {
keep = append(keep, cpu)
}
}
if len(keep) == 0 {
return nil, false
}
return keep, true
}
// parseCPUList decodes the kernel's cpulist format ("0-7,16-23", "3") into
// individual CPU IDs. Empty input yields an empty list.
func parseCPUList(s string) ([]int, error) {
if s == "" {
return nil, nil
}
var out []int
for part := range strings.SplitSeq(s, ",") {
part = strings.TrimSpace(part)
if part == "" {
continue
}
lo, hi, isRange := strings.Cut(part, "-")
a, err := strconv.Atoi(lo)
if err != nil {
return nil, fmt.Errorf("bad cpulist entry %q: %w", part, err)
}
if !isRange {
out = append(out, a)
continue
}
b, err := strconv.Atoi(hi)
if err != nil {
return nil, fmt.Errorf("bad cpulist entry %q: %w", part, err)
}
if b < a || b-a > 8192 {
return nil, fmt.Errorf("bad cpulist range %q", part)
}
for v := a; v <= b; v++ {
out = append(out, v)
}
}
return out, nil
}
func readIntFile(path string) (int, error) {
b, err := os.ReadFile(path)
if err != nil {
return 0, err
}
return strconv.Atoi(strings.TrimSpace(string(b)))
}
-163
View File
@@ -1,163 +0,0 @@
//go:build linux
package cpupick
import (
"fmt"
"os"
"path/filepath"
"slices"
"testing"
)
// fakeSysfs builds a cpuDir tree with the given per-CPU file values.
// A nil map for a file means "file absent on every CPU".
func fakeSysfs(t *testing.T, capacity, maxFreq map[int]int) string {
t.Helper()
dir := t.TempDir()
write := func(cpu int, rel string, v int) {
p := filepath.Join(dir, fmt.Sprintf("cpu%d", cpu), rel)
if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(p, fmt.Appendf(nil, "%d\n", v), 0o644); err != nil {
t.Fatal(err)
}
}
for cpu, v := range capacity {
write(cpu, "cpu_capacity", v)
}
for cpu, v := range maxFreq {
write(cpu, "cpufreq/cpuinfo_max_freq", v)
}
return dir
}
func writeCoreMask(t *testing.T, mask string) string {
t.Helper()
p := filepath.Join(t.TempDir(), "cpus")
if err := os.WriteFile(p, []byte(mask+"\n"), 0o644); err != nil {
t.Fatal(err)
}
return p
}
func TestPerfCPUsBigLittleCapacity(t *testing.T) {
// 4 big (1024) + 4 LITTLE (~290): capacity is authoritative on ARM.
dir := fakeSysfs(t, map[int]int{
0: 1024, 1: 1024, 2: 1024, 3: 1024,
4: 290, 5: 290, 6: 290, 7: 290,
}, nil)
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3, 4, 5, 6, 7})
if signal != "cpu_capacity" {
t.Fatalf("signal = %q", signal)
}
if !slices.Equal(got, []int{0, 1, 2, 3}) {
t.Errorf("got %v", got)
}
}
func TestPerfCPUsThreeTierKeepsMid(t *testing.T) {
// prime (1024) + mid (~780) + little (~280): 50% keeps prime+mid.
dir := fakeSysfs(t, map[int]int{
0: 280, 1: 280, 2: 280, 3: 280,
4: 780, 5: 780, 6: 780,
7: 1024,
}, nil)
got, _ := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3, 4, 5, 6, 7})
if !slices.Equal(got, []int{4, 5, 6, 7}) {
t.Errorf("got %v", got)
}
}
func TestPerfCPUsIntelHybridMask(t *testing.T) {
// No cpu_capacity on x86; the P-core PMU mask decides.
dir := fakeSysfs(t, nil, nil)
mask := writeCoreMask(t, "0-7")
got, signal := perfCPUsFrom(dir, mask, []int{0, 1, 2, 3, 8, 9, 10, 11})
if signal != "intel_core_pmu" {
t.Fatalf("signal = %q", signal)
}
if !slices.Equal(got, []int{0, 1, 2, 3}) {
t.Errorf("got %v", got)
}
}
func TestPerfCPUsIntelMaskDisjointFallsThrough(t *testing.T) {
// Confined to E-cores only: the mask can't help, and equal freqs below
// mean nothing else distinguishes them either -> allowed unchanged.
dir := fakeSysfs(t, nil, map[int]int{8: 4300000, 9: 4300000})
mask := writeCoreMask(t, "0-7")
got, signal := perfCPUsFrom(dir, mask, []int{8, 9})
if signal != "" || !slices.Equal(got, []int{8, 9}) {
t.Errorf("got %v signal %q", got, signal)
}
}
func TestPerfCPUsMaxFreqCompactCores(t *testing.T) {
// AMD-style compact cores: no capacity, no Intel mask; 3.3 vs 5.7 GHz.
dir := fakeSysfs(t, nil, map[int]int{
0: 5700000, 1: 5700000, 2: 3300000, 3: 3300000,
})
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3})
if signal != "max_freq" {
t.Fatalf("signal = %q", signal)
}
if !slices.Equal(got, []int{0, 1}) {
t.Errorf("got %v", got)
}
}
func TestPerfCPUsFavoredCoreSkewKept(t *testing.T) {
// Turbo Boost Max favored cores run a few percent hot; they must not
// shrink the candidate set to one or two cores.
dir := fakeSysfs(t, nil, map[int]int{
0: 5800000, 1: 5700000, 2: 5700000, 3: 5600000,
})
got, _ := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3})
if !slices.Equal(got, []int{0, 1, 2, 3}) {
t.Errorf("favored-core skew filtered CPUs: %v", got)
}
}
func TestPerfCPUsHomogeneousInconclusive(t *testing.T) {
dir := fakeSysfs(t, nil, map[int]int{0: 3000000, 1: 3000000})
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1})
if signal != "" || !slices.Equal(got, []int{0, 1}) {
t.Errorf("got %v signal %q", got, signal)
}
}
func TestPerfCPUsNoSysfs(t *testing.T) {
dir := t.TempDir()
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2})
if signal != "" || !slices.Equal(got, []int{0, 1, 2}) {
t.Errorf("got %v signal %q", got, signal)
}
}
func TestParseCPUList(t *testing.T) {
cases := []struct {
in string
want []int
wantErr bool
}{
{"0-3", []int{0, 1, 2, 3}, false},
{"0-1,16-17", []int{0, 1, 16, 17}, false},
{"5", []int{5}, false},
{"", nil, false},
{"3-1", nil, true},
{"a-b", nil, true},
{"1,x", nil, true},
}
for _, c := range cases {
got, err := parseCPUList(c.in)
if (err != nil) != c.wantErr {
t.Errorf("%q: err=%v wantErr=%v", c.in, err, c.wantErr)
continue
}
if !c.wantErr && !slices.Equal(got, c.want) {
t.Errorf("%q: got %v want %v", c.in, got, c.want)
}
}
}
-10
View File
@@ -1,10 +0,0 @@
//go:build !linux
package cpupick
// perfCPUs is Linux-only sysfs walking; elsewhere report "no distinction".
// Default already returns nil off-Linux (util.AllowedCPUs has no answer
// there), so this exists to keep the package compiling everywhere.
func perfCPUs(allowed []int) ([]int, string) {
return allowed, ""
}
-118
View File
@@ -1,118 +0,0 @@
//go:build linux
package cpupick
import (
"fmt"
"os"
"path/filepath"
"strconv"
"strings"
)
// readTopology probes the NUMA node and physical-core layout of cpus from
// sysfs. Anything sysfs won't say degrades toward flatTopology: an unknown
// node becomes node 0, an unknown core becomes a core of its own — either
// way the corresponding arrange rule becomes a no-op instead of a wrong
// answer.
func readTopology(cpus []int) topology {
return readTopologyFrom("/sys/devices/system/node", "/sys/devices/system/cpu", cpus)
}
func readTopologyFrom(nodeDir, cpuDir string, cpus []int) topology {
coreOf, zeroCore := coreGroups(cpuDir, cpus)
return topology{
nodeOf: numaNodes(nodeDir, cpus),
coreOf: coreOf,
zeroCore: zeroCore,
}
}
// numaNodes maps each cpu to its NUMA node via
// /sys/devices/system/node/nodeN/cpulist. CPUs no node claims (or no node
// dirs at all: VMs, non-NUMA kernels) land on node 0.
func numaNodes(nodeDir string, cpus []int) map[int]int {
out := make(map[int]int, len(cpus))
for _, c := range cpus {
out[c] = 0
}
entries, err := os.ReadDir(nodeDir)
if err != nil {
return out
}
want := make(map[int]bool, len(cpus))
for _, c := range cpus {
want[c] = true
}
for _, e := range entries {
id, ok := strings.CutPrefix(e.Name(), "node")
if !ok {
continue
}
n, err := strconv.Atoi(id)
if err != nil {
continue // has_cpu, possible, ... share the prefix
}
b, err := os.ReadFile(filepath.Join(nodeDir, e.Name(), "cpulist"))
if err != nil {
continue
}
list, err := parseCPUList(strings.TrimSpace(string(b)))
if err != nil {
continue
}
for _, c := range list {
if want[c] {
out[c] = n
}
}
}
return out
}
// coreGroups maps each cpu to a dense physical-core id derived from its
// (physical_package_id, core_id) pair — core_id alone repeats across
// sockets. CPUs whose topology files are unreadable get a core of their own.
// The second return is the group id of the core CPU 0 lives on, or -1 when
// that can't be determined; CPU 0's own files are consulted even when 0 is
// not a candidate, so its SMT siblings are recognized under cpusets that
// exclude CPU 0 itself.
func coreGroups(cpuDir string, cpus []int) (map[int]int, int) {
type pkgCore struct{ pkg, core int }
pairOf := func(cpu int) (pkgCore, bool) {
topoDir := filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), "topology")
pkg, err1 := readIntFile(filepath.Join(topoDir, "physical_package_id"))
core, err2 := readIntFile(filepath.Join(topoDir, "core_id"))
if err1 != nil || err2 != nil {
return pkgCore{}, false
}
return pkgCore{pkg, core}, true
}
ids := map[pkgCore]int{}
out := make(map[int]int, len(cpus))
next := 0
for _, cpu := range cpus {
k, ok := pairOf(cpu)
if !ok {
out[cpu] = next
next++
continue
}
id, ok := ids[k]
if !ok {
id = next
next++
ids[k] = id
}
out[cpu] = id
}
zeroCore := -1
if k, ok := pairOf(0); ok {
if id, ok := ids[k]; ok {
zeroCore = id
}
}
return out, zeroCore
}
-111
View File
@@ -1,111 +0,0 @@
//go:build linux
package cpupick
import (
"fmt"
"os"
"path/filepath"
"testing"
)
// fakeTopoSysfs builds nodeDir/cpuDir trees. nodes maps node id -> cpulist
// string; cores maps cpu -> (package, core) pair.
func fakeTopoSysfs(t *testing.T, nodes map[int]string, cores map[int][2]int) (string, string) {
t.Helper()
base := t.TempDir()
nodeDir := filepath.Join(base, "node")
cpuDir := filepath.Join(base, "cpu")
for n, list := range nodes {
d := filepath.Join(nodeDir, fmt.Sprintf("node%d", n))
if err := os.MkdirAll(d, 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(d, "cpulist"), []byte(list+"\n"), 0o644); err != nil {
t.Fatal(err)
}
}
for cpu, pc := range cores {
d := filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), "topology")
if err := os.MkdirAll(d, 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(d, "physical_package_id"), fmt.Appendf(nil, "%d\n", pc[0]), 0o644); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(d, "core_id"), fmt.Appendf(nil, "%d\n", pc[1]), 0o644); err != nil {
t.Fatal(err)
}
}
return nodeDir, cpuDir
}
func TestReadTopology(t *testing.T) {
// Two nodes; SMT pairs (0,4),(1,5) on node 0 and (2,6),(3,7) on node 1.
// core_id repeats across packages on purpose: the pair must disambiguate.
nodeDir, cpuDir := fakeTopoSysfs(t,
map[int]string{0: "0-1,4-5", 1: "2-3,6-7"},
map[int][2]int{
0: {0, 0}, 4: {0, 0}, 1: {0, 1}, 5: {0, 1},
2: {1, 0}, 6: {1, 0}, 3: {1, 1}, 7: {1, 1},
})
cpus := []int{0, 1, 2, 3, 4, 5, 6, 7}
topo := readTopologyFrom(nodeDir, cpuDir, cpus)
for _, c := range []int{0, 1, 4, 5} {
if topo.nodeOf[c] != 0 {
t.Errorf("cpu %d on node %d, want 0", c, topo.nodeOf[c])
}
}
for _, c := range []int{2, 3, 6, 7} {
if topo.nodeOf[c] != 1 {
t.Errorf("cpu %d on node %d, want 1", c, topo.nodeOf[c])
}
}
pairs := [][2]int{{0, 4}, {1, 5}, {2, 6}, {3, 7}}
for _, p := range pairs {
if topo.coreOf[p[0]] != topo.coreOf[p[1]] {
t.Errorf("siblings %v not grouped: %d vs %d", p, topo.coreOf[p[0]], topo.coreOf[p[1]])
}
}
if topo.coreOf[0] == topo.coreOf[2] {
t.Error("cross-package cores with equal core_id must not merge")
}
if topo.zeroCore != topo.coreOf[0] {
t.Errorf("zeroCore = %d, want %d", topo.zeroCore, topo.coreOf[0])
}
}
func TestReadTopologyZeroCoreWithoutZeroCandidate(t *testing.T) {
// CPU 0 is not a candidate (cpuset excludes it) but its sibling 4 is:
// zeroCore must still identify their shared core.
nodeDir, cpuDir := fakeTopoSysfs(t,
map[int]string{0: "0-7"},
map[int][2]int{0: {0, 0}, 4: {0, 0}, 1: {0, 1}, 5: {0, 1}})
topo := readTopologyFrom(nodeDir, cpuDir, []int{1, 4, 5})
if topo.zeroCore < 0 || topo.coreOf[4] != topo.zeroCore {
t.Errorf("zeroCore = %d, coreOf[4] = %d; sibling of CPU 0 not identified", topo.zeroCore, topo.coreOf[4])
}
if topo.coreOf[1] == topo.zeroCore {
t.Error("cpu 1 wrongly grouped with CPU 0's core")
}
}
func TestReadTopologyMissingSysfs(t *testing.T) {
base := t.TempDir()
cpus := []int{0, 1, 2}
topo := readTopologyFrom(filepath.Join(base, "nope"), filepath.Join(base, "also-nope"), cpus)
seen := map[int]bool{}
for _, c := range cpus {
if topo.nodeOf[c] != 0 {
t.Errorf("cpu %d node = %d, want 0", c, topo.nodeOf[c])
}
if seen[topo.coreOf[c]] {
t.Errorf("cpu %d shares a fallback core group", c)
}
seen[topo.coreOf[c]] = true
}
if topo.zeroCore != -1 {
t.Errorf("zeroCore = %d, want -1 when unknown", topo.zeroCore)
}
}
-10
View File
@@ -1,10 +0,0 @@
//go:build !linux
package cpupick
// readTopology has no sysfs to consult off Linux; the flat stand-in makes
// arrange's NUMA and SMT rules no-ops. Default is already nil off Linux
// (util.AllowedCPUs has no answer there) — this keeps the package compiling.
func readTopology(cpus []int) topology {
return flatTopology(cpus)
}
+7 -16
View File
@@ -97,7 +97,8 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
newAddr := getDnsServerAddr(c)
d.serverMu.Lock()
running := d.server != nil
running := d.server
runningStarted := d.started
sameAddr := d.addr == newAddr
d.addr = newAddr
d.enabled.Store(enabled)
@@ -111,7 +112,7 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
}
if !enabled {
if running {
if running != nil {
d.Stop()
}
// Drop any records that accumulated while enabled; a later re-enable
@@ -120,12 +121,12 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
return nil
}
if !running {
if running == nil {
// Was disabled (or never started); bring it up now.
go d.Start()
} else if !sameAddr {
// Stop clears the slot before shutting down, otherwise the Start below can find the dying server and refuse
d.Stop()
d.shutdownServer(running, runningStarted, "reload")
// Old Start goroutine has now exited; bring up a fresh listener on the new address.
go d.Start()
}
@@ -161,9 +162,7 @@ func (d *dnsServer) Start() {
started := make(chan struct{})
d.serverMu.Lock()
// Re-check enabled under the lock, a disable that raced our check above snapshots the slot under it too.
// Two reloads in quick succession can both spawn a Start, the loser would orphan the live listener past Stop
if d.ctx.Err() != nil || d.server != nil || !d.enabled.Load() {
if d.ctx.Err() != nil {
d.serverMu.Unlock()
return
}
@@ -201,14 +200,6 @@ func (d *dnsServer) Start() {
close(started)
}
// Release our slot, unless a reload already replaced us, so a dead listener can't block a future Start
d.serverMu.Lock()
if d.server == server {
d.server = nil
d.started = nil
}
d.serverMu.Unlock()
if err != nil {
d.l.Warn("Failed to run the DNS responder", "error", err)
}
+4 -206
View File
@@ -194,51 +194,14 @@ func TestDnsServer_reload_initial_serveDnsWithoutLighthouse(t *testing.T) {
}
func TestDnsServer_reload_sameAddr_noOp(t *testing.T) {
port := freeUDPPort(t)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", port, true, true)
setDnsConfig(c, "127.0.0.1", "0", true, true)
require.NoError(t, ds.reload(c, true))
go ds.Start()
waitForBind(t, ds)
ds.serverMu.Lock()
before := ds.server
ds.serverMu.Unlock()
require.NotNil(t, before)
// Same address, so the running listener must be left alone rather than rebuilt under live queries
// No server running yet, no addr change. Reload should not spawn anything.
require.NoError(t, ds.reload(c, false))
assert.True(t, ds.enabled.Load())
ds.serverMu.Lock()
after := ds.server
ds.serverMu.Unlock()
assert.Same(t, before, after, "a same-address reload must not restart the listener")
ds.Stop()
}
// The branch the old sameAddr test was accidentally hitting: enabled with nothing running means reload starts it.
func TestDnsServer_reload_whenNotRunning_starts(t *testing.T) {
port := freeUDPPort(t)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", port, true, true)
// initial only records config, it never starts anything
require.NoError(t, ds.reload(c, true))
ds.serverMu.Lock()
assert.Nil(t, ds.server, "the initial reload must not start a listener")
ds.serverMu.Unlock()
require.NoError(t, ds.reload(c, false))
waitForBind(t, ds)
ds.serverMu.Lock()
assert.NotNil(t, ds.server, "a reload with nothing running should bring DNS up")
ds.serverMu.Unlock()
ds.Stop()
assert.Nil(t, ds.server)
}
func TestDnsServer_StartStop_lifecycle(t *testing.T) {
@@ -464,168 +427,3 @@ func waitFor(t *testing.T, cond func() bool) {
}
t.Fatal("timed out waiting for condition")
}
// Two reloads in quick succession, or a HUP before Control.Start, can race two Starts at the same listener.
func TestDnsServer_Start_isIdempotent(t *testing.T) {
port := freeUDPPort(t)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", port, true, true)
require.NoError(t, ds.reload(c, true))
go ds.Start()
waitForBind(t, ds)
ds.serverMu.Lock()
first := ds.server
ds.serverMu.Unlock()
require.NotNil(t, first)
// If the second Start replaces the tracked server, Stop kills the wrong one and the port leaks
done := make(chan struct{})
go func() {
ds.Start()
close(done)
}()
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("second Start never returned")
}
ds.serverMu.Lock()
second := ds.server
ds.serverMu.Unlock()
assert.Same(t, first, second, "a second Start must not replace the running server")
// The real proof, after Stop the port must actually be free
ds.Stop()
waitFor(t, func() bool {
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
if err != nil {
return false
}
_ = pc.Close()
return true
})
}
// An address change must actually end up listening on the new port. Start's guard refuses when a server is already
// installed, so reload has to clear the slot before shutting the old one down.
func TestDnsServer_reload_addrChange_restarts(t *testing.T) {
first := freeUDPPort(t)
second := freeUDPPort(t)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", first, true, true)
require.NoError(t, ds.reload(c, true))
go ds.Start()
waitForBind(t, ds)
// Cycle a few times, the failure this guards against depends on which goroutine wins serverMu
for i := range 8 {
want := second
if i%2 == 1 {
want = first
}
setDnsConfig(c, "127.0.0.1", want, true, true)
require.NoError(t, ds.reload(c, false))
waitForBind(t, ds)
ds.serverMu.Lock()
srv := ds.server
ds.serverMu.Unlock()
require.NotNil(t, srv, "reload left DNS down instead of restarting it")
require.Equal(t, "127.0.0.1:"+want, srv.Addr, "reload should be serving the new address")
}
// Land back on second so the port assertions below are meaningful
setDnsConfig(c, "127.0.0.1", second, true, true)
require.NoError(t, ds.reload(c, false))
waitForBind(t, ds)
// The old port must be released and the new one actually held
waitFor(t, func() bool {
pc, err := net.ListenPacket("udp", "127.0.0.1:"+first)
if err != nil {
return false
}
_ = pc.Close()
return true
})
_, err := net.ListenPacket("udp", "127.0.0.1:"+second)
require.Error(t, err, "the new address should be bound by the DNS responder")
ds.Stop()
}
// A listener that dies on its own must release the slot, or a later same-addr reload sees it as running and no-ops.
func TestDnsServer_Start_bindFailure_releasesSlot(t *testing.T) {
port := freeUDPPort(t)
blocker, err := net.ListenPacket("udp", "127.0.0.1:"+port)
require.NoError(t, err)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", port, true, true)
require.NoError(t, ds.reload(c, true))
ds.Start() // returns once the bind fails
ds.serverMu.Lock()
assert.Nil(t, ds.server, "a listener that failed to bind must not stay parked in the slot")
ds.serverMu.Unlock()
// With the slot released, a reload can retry once the port frees up
require.NoError(t, blocker.Close())
require.NoError(t, ds.reload(c, false))
waitForBind(t, ds)
ds.serverMu.Lock()
assert.NotNil(t, ds.server, "a same-addr reload should retry after a failed bind")
ds.serverMu.Unlock()
ds.Stop()
}
// A disable that lands while Start is between its unlocked check and the guard must not leave a listener behind.
func TestDnsServer_Start_refusesWhenDisabledUnderLock(t *testing.T) {
port := freeUDPPort(t)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", port, true, true)
require.NoError(t, ds.reload(c, true))
require.True(t, ds.enabled.Load())
// Holding serverMu parks Start on the lock, the only way to land the disable in that window on purpose
ds.serverMu.Lock()
done := make(chan struct{})
go func() {
ds.Start()
close(done)
}()
select {
case <-done:
ds.serverMu.Unlock()
t.Fatal("Start returned early, the test never exercised the window")
case <-time.After(time.Millisecond * 100):
}
// The disable reload's critical section. It sees nothing running, so it never calls Stop.
ds.enabled.Store(false)
ds.serverMu.Unlock()
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("Start never returned")
}
ds.serverMu.Lock()
assert.Nil(t, ds.server, "Start must not install a listener a disable already cancelled")
ds.serverMu.Unlock()
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
require.NoError(t, err, "an orphaned listener is still holding the port")
_ = pc.Close()
}
+30 -98
View File
@@ -405,7 +405,7 @@ func TestStage1Race(t *testing.T) {
r.Log("Spin until connection manager tears down a tunnel")
for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 {
for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 {
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
t.Log("Connection manager hasn't ticked yet")
time.Sleep(time.Second)
@@ -453,11 +453,9 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
r.Log("Nuke my hostmap")
myHostmap := myControl.GetHostmap()
myHostmap.Lock()
myHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
myHostmap.Indexes = map[uint32]*nebula.HostInfo{}
myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
myHostmap.Unlock()
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again")))
p = r.RouteForAllUntilTxTun(theirControl)
@@ -467,10 +465,10 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
r.Log("Wait for the dead index to go away")
start := theirControl.GetHostmapIndexCount()
start := len(theirControl.GetHostmap().Indexes)
for {
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
if theirControl.GetHostmapIndexCount() < start {
if len(theirControl.GetHostmap().Indexes) < start {
break
}
time.Sleep(time.Second)
@@ -506,11 +504,9 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
r.Log("Nuke my hostmap")
theirHostmap := theirControl.GetHostmap()
theirHostmap.Lock()
theirHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
theirHostmap.Indexes = map[uint32]*nebula.HostInfo{}
theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
theirHostmap.Unlock()
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again")))
p = r.RouteForAllUntilTxTun(myControl)
@@ -521,10 +517,10 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
r.Log("Wait for the dead index to go away")
start := myControl.GetHostmapIndexCount()
start := len(myControl.GetHostmap().Indexes)
for {
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
if myControl.GetHostmapIndexCount() < start {
if len(myControl.GetHostmap().Indexes) < start {
break
}
time.Sleep(time.Second)
@@ -632,10 +628,10 @@ func TestReestablishRelays(t *testing.T) {
r.Log("Close the tunnel")
relayControl.CloseTunnel(theirVpnIpNet[0].Addr(), true)
start := myControl.GetHostmapIndexCount()
curIndexes := myControl.GetHostmapIndexCount()
start := len(myControl.GetHostmap().Indexes)
curIndexes := len(myControl.GetHostmap().Indexes)
for curIndexes >= start {
curIndexes = myControl.GetHostmapIndexCount()
curIndexes = len(myControl.GetHostmap().Indexes)
r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes)
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail")))
@@ -725,70 +721,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
@@ -887,18 +819,18 @@ func TestStage1RaceRelays2(t *testing.T) {
t.Log("Wait until we remove extra tunnels")
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
myControl.GetHostmapIndexCount(),
theirControl.GetHostmapIndexCount(),
relayControl.GetHostmapIndexCount(),
len(myControl.GetHostmap().Indexes),
len(theirControl.GetHostmap().Indexes),
len(relayControl.GetHostmap().Indexes),
)
hostInfos := myControl.GetHostmapIndexCount() + theirControl.GetHostmapIndexCount() + relayControl.GetHostmapIndexCount()
hostInfos := len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes)
retries := 60
for hostInfos > 6 && retries > 0 {
hostInfos = myControl.GetHostmapIndexCount() + theirControl.GetHostmapIndexCount() + relayControl.GetHostmapIndexCount()
hostInfos = len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes)
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
myControl.GetHostmapIndexCount(),
theirControl.GetHostmapIndexCount(),
relayControl.GetHostmapIndexCount(),
len(myControl.GetHostmap().Indexes),
len(theirControl.GetHostmap().Indexes),
len(relayControl.GetHostmap().Indexes),
)
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
t.Log("Connection manager hasn't ticked yet")
@@ -992,24 +924,24 @@ func TestRehandshakingRelays(t *testing.T) {
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl)
// We should have two hostinfos on all sides
for myControl.GetHostmapIndexCount() != 2 {
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", myControl.GetHostmapIndexCount())
for len(myControl.GetHostmap().Indexes) != 2 {
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(myControl.GetHostmap().Indexes))
r.Log("Assert the relay tunnel still works")
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.Log("yupitdoes")
time.Sleep(time.Second)
}
t.Logf("myControl hostinfos got cleaned up!")
for theirControl.GetHostmapIndexCount() != 2 {
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", theirControl.GetHostmapIndexCount())
for len(theirControl.GetHostmap().Indexes) != 2 {
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(theirControl.GetHostmap().Indexes))
r.Log("Assert the relay tunnel still works")
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.Log("yupitdoes")
time.Sleep(time.Second)
}
t.Logf("theirControl hostinfos got cleaned up!")
for relayControl.GetHostmapIndexCount() != 2 {
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", relayControl.GetHostmapIndexCount())
for len(relayControl.GetHostmap().Indexes) != 2 {
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(relayControl.GetHostmap().Indexes))
r.Log("Assert the relay tunnel still works")
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.Log("yupitdoes")
@@ -1097,24 +1029,24 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl)
// We should have two hostinfos on all sides
for myControl.GetHostmapIndexCount() != 2 {
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", myControl.GetHostmapIndexCount())
for len(myControl.GetHostmap().Indexes) != 2 {
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(myControl.GetHostmap().Indexes))
r.Log("Assert the relay tunnel still works")
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.Log("yupitdoes")
time.Sleep(time.Second)
}
t.Logf("myControl hostinfos got cleaned up!")
for theirControl.GetHostmapIndexCount() != 2 {
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", theirControl.GetHostmapIndexCount())
for len(theirControl.GetHostmap().Indexes) != 2 {
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(theirControl.GetHostmap().Indexes))
r.Log("Assert the relay tunnel still works")
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.Log("yupitdoes")
time.Sleep(time.Second)
}
t.Logf("theirControl hostinfos got cleaned up!")
for relayControl.GetHostmapIndexCount() != 2 {
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", relayControl.GetHostmapIndexCount())
for len(relayControl.GetHostmap().Indexes) != 2 {
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(relayControl.GetHostmap().Indexes))
r.Log("Assert the relay tunnel still works")
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.Log("yupitdoes")
@@ -1191,7 +1123,7 @@ func TestRehandshaking(t *testing.T) {
theirConfig.ReloadConfigString(string(rc))
r.Log("Spin until there is only 1 tunnel")
for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 {
for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 {
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
t.Log("Connection manager hasn't ticked yet")
time.Sleep(time.Second)
@@ -1291,7 +1223,7 @@ func TestRehandshakingLoser(t *testing.T) {
myConfig.ReloadConfigString(string(rc))
r.Log("Spin until there is only 1 tunnel")
for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 {
for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 {
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
t.Log("Connection manager hasn't ticked yet")
time.Sleep(time.Second)
+4 -2
View File
@@ -4,13 +4,15 @@
package e2e
import (
"log/slog"
"io"
"net/netip"
"os"
"strings"
"testing"
"time"
"log/slog"
"dario.cat/mergo"
"github.com/google/gopacket"
"github.com/google/gopacket/layers"
@@ -380,7 +382,7 @@ func getAddrs(ns []netip.Prefix) []netip.Addr {
func NewTestLogger() *slog.Logger {
v := os.Getenv("TEST_LOGS")
if v == "" {
return slog.New(slog.DiscardHandler)
return slog.New(slog.NewTextHandler(io.Discard, nil))
}
level := slog.LevelInfo
-225
View File
@@ -1,225 +0,0 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
import (
"net/netip"
"testing"
"time"
"github.com/slackhq/nebula"
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/cert_test"
"github.com/slackhq/nebula/e2e/router"
"github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/udp"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// reportedAddrs is what the lighthouse would hand a peer asking where vpnAddr is.
func reportedAddrs(t *testing.T, lh *nebula.Control, vpnAddr netip.Addr) []netip.AddrPort {
t.Helper()
cm := lh.QueryLighthouse(vpnAddr)
if cm == nil {
return nil
}
var out []netip.AddrPort
for _, c := range *cm {
out = append(out, c.Reported...)
out = append(out, c.Learned...)
}
return out
}
// waitForLighthouseMsg routes until a lighthouse message lands on lh, or gives up. Reports whether one arrived.
func waitForLighthouseMsg(t *testing.T, r *router.R, lh *nebula.Control, wait time.Duration) bool {
t.Helper()
h := &header.H{}
return r.RouteForAllExitFuncOrTimeout(wait, func(p *udp.Packet, c *nebula.Control) router.ExitType {
if c != lh {
return router.KeepRouting
}
// Punches are a single byte and never parse, they are just not what we are after
if err := h.Parse(p.Data); err != nil {
return router.KeepRouting
}
if h.Type == header.LightHouse {
return router.RouteAndExit
}
return router.KeepRouting
})
}
// A laptop that changes networks has to tell the lighthouse promptly, otherwise the lighthouse keeps handing peers
// the old address and their punches land nowhere. On a long lighthouse interval the only thing that closes that
// window is the rebind, which on darwin the network change monitor drives. The e2e build compiles the monitor out,
// so we call RebindUDPServer directly, which is the same thing the monitor does.
func TestRebindSendsLighthouseUpdate(t *testing.T) {
t.Parallel()
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
"lighthouse": m{"am_lighthouse": true},
})
// 600s interval, so nothing scheduled can send an update during this test. A rebind is the only thing that can.
myControl, _, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
"lighthouse": m{
"hosts": []any{lhVpnIpNet[0].Addr().String()},
"interval": 600,
},
"static_host_map": m{
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
},
})
r := router.NewR(t, lhControl, myControl)
defer r.RenderFlow()
lhControl.Start()
myControl.Start()
// Let the startup registration finish, then clear everything it left behind
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
r.RouteFor(time.Millisecond * 400)
// Nothing should be talking to the lighthouse on its own now
require.False(t, waitForLighthouseMsg(t, r, lhControl, time.Millisecond*200),
"nothing should reach the lighthouse before the rebind")
myControl.RebindUDPServer()
assert.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5),
"a rebind should push an update to the lighthouse rather than waiting out the interval")
lhControl.Stop()
myControl.Stop()
}
// The other half of a rebind: every live tunnel requeries the lighthouse on its next send. That query is what makes
// the lighthouse tell the peer to punch toward our new address, which is the part that actually revives a tunnel
// whose remote NAT state died while we were on a different network.
func TestRebindRequeriesPeersOnNextSend(t *testing.T) {
t.Parallel()
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
"lighthouse": m{"am_lighthouse": true},
})
lhCfg := m{
"lighthouse": m{
"hosts": []any{lhVpnIpNet[0].Addr().String()},
"interval": 600,
// Without this the peers advertise this machine's real addresses and then try to punch at them,
// which the router has no route for.
"local_allow_list": m{
"10.0.0.0/24": true,
"::/0": false,
},
},
"static_host_map": m{
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
},
}
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", lhCfg)
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.3/24", lhCfg)
r := router.NewR(t, lhControl, myControl, theirControl)
defer r.RenderFlow()
lhControl.Start()
myControl.Start()
theirControl.Start()
r.RouteFor(time.Millisecond * 500)
// Point the peers at each other directly, this test is about the rebind and not about lighthouse discovery
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("initial")))
r.RouteFor(time.Second)
require.NotNil(t, myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false), "expected a tunnel to them")
r.RouteFor(time.Millisecond * 300)
// Assert on what the peer sees rather than on lighthouse traffic. A query for them makes the lighthouse send
// them a punch notification, which is the whole point. Our own update to the lighthouse sends them nothing,
// so this cannot be satisfied by the update the rebind itself pushes.
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("quiet")))
require.False(t, waitForLighthouseMsg(t, r, theirControl, time.Millisecond*300),
"an ordinary send should not requery the lighthouse")
myControl.RebindUDPServer()
r.RouteFor(time.Millisecond * 300) // let the update the rebind itself sends pass by
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("after rebind")))
assert.True(t, waitForLighthouseMsg(t, r, theirControl, time.Second*5),
"the first send after a rebind should requery the lighthouse, which then tells the peer to punch at us")
lhControl.Stop()
myControl.Stop()
theirControl.Stop()
}
// The scenario this whole thing exists for: a laptop sleeps at the office and wakes up at home on a new address.
// Until it tells the lighthouse, the lighthouse keeps handing peers the office address, so their punches land
// nowhere and the tunnel stays dead. On a long interval the rebind is the only thing that closes that window.
func TestRebindAdvertisesNewAddressAfterMove(t *testing.T) {
t.Parallel()
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
"lighthouse": m{"am_lighthouse": true},
})
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
"lighthouse": m{
"hosts": []any{lhVpnIpNet[0].Addr().String()},
"interval": 600,
},
"static_host_map": m{
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
},
})
// Advertise wherever we currently are rather than this machine's real NICs, read fresh each time so a move
// is picked up.
myControl.SetLocalAddrsFn(func(*nebula.LocalAllowList) []netip.Addr {
return []netip.Addr{myControl.GetUDPAddr().Addr()}
})
r := router.NewR(t, lhControl, myControl)
defer r.RenderFlow()
lhControl.Start()
myControl.Start()
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
r.RouteFor(time.Millisecond * 400)
require.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), myUdpAddr,
"the lighthouse should know the address we started on")
// Wake up somewhere else
newAddr := netip.MustParseAddrPort("10.0.0.99:4242")
myControl.SetUDPAddr(newAddr)
r.AddRoute(newAddr.Addr(), newAddr.Port(), myControl)
// Nothing has told the lighthouse, and with interval 600 nothing scheduled will
r.RouteFor(time.Millisecond * 400)
require.NotContains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
"the lighthouse should still be handing out the old address before the rebind")
myControl.RebindUDPServer()
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an update after the rebind")
r.RouteFor(time.Millisecond * 400)
assert.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
"after the rebind the lighthouse should hand peers our new address")
lhControl.Stop()
myControl.Stop()
}
-136
View File
@@ -1,136 +0,0 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
import (
"testing"
"time"
"github.com/slackhq/nebula"
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/cert_test"
"github.com/slackhq/nebula/e2e/router"
"github.com/slackhq/nebula/udp"
)
// TestRecoveryTiming measures how long a tunnel takes to come back after the peer stops accepting our traffic,
// which is what a laptop waking on a new network looks like from the peer's side: its NAT has no state for where
// we are now, so everything we send disappears.
//
// It is a measurement, not a pass/fail assertion. Recovery is timed to the moment the peer punches back at us,
// since that is when its NAT opens and the tunnel is usable again.
//
// go test -tags e2e_testing -v -run TestRecoveryTiming ./e2e/
func TestRecoveryTiming(t *testing.T) {
for _, tc := range []struct {
name string
rebind bool
}{
{"no trigger", false},
{"rebind counter", true},
} {
t.Run(tc.name, func(t *testing.T) {
d, lost := measureRecovery(t, tc.rebind)
t.Logf("RESULT %-16s recovered in %-9v (%d packets lost)", tc.name, d.Round(time.Millisecond), lost)
})
}
}
// measureRecovery returns how long until the peer punched back, and how many of our packets died meanwhile. When
// rebind is true we call RebindUDPServer once the tunnel goes dark, which is what the darwin network change
// monitor does and what iOS has always done. When false, nothing tells nebula anything is wrong.
func measureRecovery(t *testing.T, rebind bool) (time.Duration, int) {
t.Helper()
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
"lighthouse": m{"am_lighthouse": true},
})
peerCfg := m{
"lighthouse": m{
"hosts": []any{lhVpnIpNet[0].Addr().String()},
"interval": 600,
"local_allow_list": m{
"10.0.0.0/24": true,
"::/0": false,
},
},
"static_host_map": m{
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
},
}
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", peerCfg)
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.3/24", peerCfg)
r := router.NewR(t, lhControl, myControl, theirControl)
defer r.RenderFlow()
defer func() {
lhControl.Stop()
myControl.Stop()
theirControl.Stop()
}()
lhControl.Start()
myControl.Start()
theirControl.Start()
r.RouteFor(time.Millisecond * 500)
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("establish")))
r.RouteFor(time.Second)
if myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false) == nil {
t.Fatal("failed to establish the tunnel we are measuring")
}
r.RouteFor(time.Millisecond * 500)
// From here the peer's NAT has no state for us, everything we send it disappears
start := time.Now()
blackholed := 0
var recovered time.Duration
if rebind {
myControl.RebindUDPServer()
}
// Keep the tun busy the way someone retrying a stalled connection would
stop := make(chan struct{})
defer close(stop)
go func() {
tick := time.NewTicker(time.Millisecond * 200)
defer tick.Stop()
for {
select {
case <-stop:
return
case <-tick.C:
myControl.InjectTunPacket(BuildTunUDPPacket(
theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("retry")))
}
}
}()
r.RouteForAllExitFuncOrTimeout(time.Second*30, func(p *udp.Packet, c *nebula.Control) router.ExitType {
if c == theirControl && p.From == myControl.GetUDPAddr() {
blackholed++
return router.Drop
}
// The peer reaching us directly is the moment its NAT opened, whether that is a punch or a handshake
if c == myControl && p.From == theirUdpAddr {
recovered = time.Since(start)
return router.RouteAndExit
}
return router.KeepRouting
})
if recovered == 0 {
t.Fatalf("no recovery within 30s (%d packets blackholed)", blackholed)
}
return recovered, blackholed
}
+30 -149
View File
@@ -114,28 +114,6 @@ type packet struct {
packet *udp.Packet
tun bool // a packet pulled off a tun device
rx bool // the packet was received by a udp device
// h is the nebula header, parsed once when the packet is recorded. parseErr says why there isn't one, which
// the flow log reports rather than hiding. Punchy sends a single byte, so an unparseable packet is normal.
h header.H
parseErr error
}
// fromAddr and toAddr are the addresses this packet actually travelled between. Reading them off the control
// instead would misreport the whole history once a test moves a node. Tun packets are synthesized without
// addresses, so they fall back to the control.
func (p *packet) fromAddr() netip.AddrPort {
if p.tun || !p.packet.From.IsValid() {
return p.from.GetUDPAddr()
}
return p.packet.From
}
func (p *packet) toAddr() netip.AddrPort {
if p.tun || !p.packet.To.IsValid() {
return p.to.GetUDPAddr()
}
return p.packet.To
}
func (p *packet) WasReceived() {
@@ -153,9 +131,6 @@ const (
ExitNow ExitType = 1
// RouteAndExit routes this packet and exits immediately afterwards
RouteAndExit ExitType = 2
// Drop discards this packet without delivering it and keeps routing. Use it to simulate a blackhole, such as
// a restrictive NAT refusing traffic from an address it has not seen.
Drop ExitType = 3
)
type ExitFunc func(packet *udp.Packet, receiver *nebula.Control) ExitType
@@ -166,9 +141,7 @@ type ExitFunc func(packet *udp.Packet, receiver *nebula.Control) ExitType
func NewR(t testing.TB, controls ...*nebula.Control) *R {
ctx, cancel := context.WithCancel(context.Background())
// t.Name() contains a slash for subtests, so the flow log can land in a nested directory
fn := filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name()))
if err := os.MkdirAll(filepath.Dir(fn), 0755); err != nil {
if err := os.MkdirAll("mermaid", 0755); err != nil {
panic(err)
}
@@ -179,7 +152,7 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
outNat: make(map[outNatKey]netip.AddrPort),
flow: []flowEntry{},
ignoreFlows: []ignoreFlow{},
fn: fn,
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())),
t: t,
cancelRender: cancel,
}
@@ -276,7 +249,7 @@ func (r *R) renderFlow() {
continue
}
addr := e.packet.fromAddr()
addr := e.packet.from.GetUDPAddr()
if _, ok := participants[addr]; ok {
continue
}
@@ -295,6 +268,7 @@ func (r *R) renderFlow() {
}
// Print packets
h := &header.H{}
for _, e := range r.flow {
if e.packet == nil {
//fmt.Fprintf(f, " note over %s: %s\n", strings.Join(participantsVals, ", "), e.note)
@@ -306,22 +280,21 @@ func (r *R) renderFlow() {
fmt.Fprintln(f, r.formatUdpPacket(p))
} else {
if err := h.Parse(p.packet.Data); err != nil {
panic(err)
}
line := "--x"
if p.rx {
line = "->>"
}
detail := fmt.Sprintf("%s(%s), index %v, counter: %v",
p.h.TypeName(), p.h.SubTypeName(), p.h.RemoteIndex, p.h.MessageCounter)
if p.parseErr != nil {
detail = fmt.Sprintf("unparsed, %v (%d bytes)", p.parseErr, len(p.packet.Data))
}
fmt.Fprintf(f, " %s%s%s: %s\n",
normalizeName(p.fromAddr().String()),
fmt.Fprintf(f,
" %s%s%s: %s(%s), index %v, counter: %v\n",
normalizeName(p.from.GetUDPAddr().String()),
line,
normalizeName(p.toAddr().String()),
detail,
normalizeName(p.to.GetUDPAddr().String()),
h.TypeName(), h.SubTypeName(), h.RemoteIndex, h.MessageCounter,
)
}
}
@@ -435,34 +408,29 @@ func (r *R) unlockedInjectFlow(from, to *nebula.Control, p *udp.Packet, tun bool
r.renderHostmaps(fmt.Sprintf("Packet %v", len(r.flow)))
var h header.H
var parseErr error
if !tun {
parseErr = h.Parse(p.Data)
}
// Decide before copying, the copy comes from a freelist and an ignored packet would never be released
for _, i := range r.ignoreFlows {
if tun {
if i.tun.HasValue && i.tun.IsTrue {
return nil
}
continue
if len(r.ignoreFlows) > 0 {
var h header.H
err := h.Parse(p.Data)
if err != nil {
panic(err)
}
// A packet we could not parse has no type to match against, so no rule can ignore it
if parseErr == nil && i.messageType == h.Type && i.subType == h.Subtype {
return nil
for _, i := range r.ignoreFlows {
if !tun {
if i.messageType == h.Type && i.subType == h.Subtype {
return nil
}
} else if i.tun.HasValue && i.tun.IsTrue {
return nil
}
}
}
fp := &packet{
from: from,
to: to,
packet: p.Copy(),
tun: tun,
h: h,
parseErr: parseErr,
from: from,
to: to,
packet: p.Copy(),
tun: tun,
}
r.flow = append(r.flow, flowEntry{packet: fp})
@@ -692,10 +660,6 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
p.Release()
return
case Drop:
// Record it so the flow log shows the attempt, but never hand it to the receiver
r.unlockedInjectFlow(sender, receiver, p, false)
case KeepRouting:
fp := r.unlockedInjectFlow(sender, receiver, p, false)
receiver.InjectUDPPacket(p)
@@ -726,85 +690,6 @@ func (r *R) RouteUntilAfterMsgType(sender *nebula.Control, msgType header.Messag
})
}
// RouteFor routes everything that shows up for the given duration and then returns. Use it to let a test settle
// deterministically rather than sleeping and hoping: a single FlushAll races a completing handshake, which queues
// more packets right behind it.
func (r *R) RouteFor(d time.Duration) {
r.RouteForAllExitFuncOrTimeout(d, func(*udp.Packet, *nebula.Control) ExitType {
return KeepRouting
})
}
// RouteForAllExitFuncOrTimeout is RouteForAllExitFunc with a deadline, reporting whether whatDo asked to exit
// before time ran out. The unbounded version blocks forever on a quiet network, so this is what a test needs to
// assert that something does NOT happen, or to route for a fixed settling period.
func (r *R) RouteForAllExitFuncOrTimeout(timeout time.Duration, whatDo ExitFunc) bool {
sc := make([]reflect.SelectCase, 0, len(r.controls)+1)
cm := make([]*nebula.Control, 0, len(r.controls))
for _, c := range r.controls {
sc = append(sc, reflect.SelectCase{
Dir: reflect.SelectRecv,
Chan: reflect.ValueOf(c.GetUDPTxChan()),
Send: reflect.Value{},
})
cm = append(cm, c)
}
timer := time.NewTimer(timeout)
defer timer.Stop()
sc = append(sc, reflect.SelectCase{
Dir: reflect.SelectRecv,
Chan: reflect.ValueOf(timer.C),
Send: reflect.Value{},
})
for {
x, rx, _ := reflect.Select(sc)
if x == len(cm) {
return false
}
r.Lock()
p := rx.Interface().(*udp.Packet)
receiver := r.getControl(cm[x].GetUDPAddr(), p.To, p)
if receiver == nil {
r.Unlock()
panic("Can't RouteForAllExitFuncOrTimeout for host: " + p.To.String())
}
e := whatDo(p, receiver)
switch e {
case ExitNow:
r.Unlock()
p.Release()
return true
case RouteAndExit:
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
receiver.InjectUDPPacket(p)
fp.WasReceived()
r.Unlock()
p.Release()
return true
case Drop:
// Record it so the flow log shows the attempt, but never hand it to the receiver
r.unlockedInjectFlow(cm[x], receiver, p, false)
case KeepRouting:
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
receiver.InjectUDPPacket(p)
fp.WasReceived()
default:
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
}
r.Unlock()
p.Release()
}
}
func (r *R) RouteForAllUntilAfterMsgTypeTo(receiver *nebula.Control, msgType header.MessageType, subType header.MessageSubType) {
h := &header.H{}
r.RouteForAllExitFunc(func(p *udp.Packet, r *nebula.Control) ExitType {
@@ -897,10 +782,6 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
p.Release()
return
case Drop:
// Record it so the flow log shows the attempt, but never hand it to the receiver
r.unlockedInjectFlow(cm[x], receiver, p, false)
case KeepRouting:
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
receiver.InjectUDPPacket(p)
+6 -6
View File
@@ -43,8 +43,8 @@ func TestDropInactiveTunnels(t *testing.T) {
r.Log("Go inactive and wait for the tunnels to get dropped")
waitStart := time.Now()
for {
myIndexes := myControl.GetHostmapIndexCount()
theirIndexes := theirControl.GetHostmapIndexCount()
myIndexes := len(myControl.GetHostmap().Indexes)
theirIndexes := len(theirControl.GetHostmap().Indexes)
if myIndexes == 0 && theirIndexes == 0 {
break
}
@@ -493,8 +493,8 @@ func TestCloseTunnelAuthenticated(t *testing.T) {
waitStart := time.Now()
for {
myIndexes := myControl.GetHostmapIndexCount()
theirIndexes := theirControl.GetHostmapIndexCount()
myIndexes := len(myControl.GetHostmap().Indexes)
theirIndexes := len(theirControl.GetHostmap().Indexes)
if myIndexes == 0 && theirIndexes == 0 {
break
}
@@ -548,8 +548,8 @@ func TestCloseTunnelAuthenticated(t *testing.T) {
r.Log("Injected bogus close tunnel. Let's see!")
waitStart = time.Now()
for {
myIndexes := myControl.GetHostmapIndexCount()
theirIndexes := theirControl.GetHostmapIndexCount()
myIndexes := len(myControl.GetHostmap().Indexes)
theirIndexes := len(theirControl.GetHostmap().Indexes)
if myIndexes == 0 {
t.Fatal("myIndexes should not be 0")
}
+4 -30
View File
@@ -131,9 +131,6 @@ listen:
port: 4242
# Sets the max number of packets to pull from the kernel for each syscall (under systems that support recvmmsg)
# default is 64, does not support reload
# Note: on Linux with UDP GRO (kernel 5.10+), each receive slot is sized for a full 64KiB coalesced
# superpacket, so the receive scratch is batch * 64KiB per listening socket (~4MiB per routine at the
# default of 64). Lower this to trade peak per-syscall throughput for memory on constrained hosts.
#batch: 64
# Configure socket buffers for the udp side (outside), leave unset to use the system defaults. Values will be doubled by the kernel
# Default is net.core.rmem_default and net.core.wmem_default (/proc/sys/net/core/rmem_default and /proc/sys/net/core/rmem_default)
@@ -149,14 +146,6 @@ listen:
# Default true; set to false to leave WDF in charge of inbound decisions on the listener port. Not reloadable.
#windows_bypass_wdf: true
# On macOS only
# macOS scopes the udp socket to the interface it was created on, so moving between networks (wifi to wired,
# office to home) leaves Nebula sending out an interface that no longer has a route. When true, Nebula watches
# the routing socket and rebinds the listener once the change settles.
# iOS does not use this, the host app drives the same rebind itself.
# Default true. Not reloadable.
#rebind_on_network_change: true
# By default, Nebula replies to packets it has no tunnel for with a "recv_error" packet. This packet helps speed up reconnection
# in the case that Nebula on either side did not shut down cleanly. This response can be abused as a way to discover if Nebula is running
# on a host though. This option lets you configure if you want to send "recv_error" packets always, never, or only to private network remotes.
@@ -253,6 +242,10 @@ tun:
# When tun is disabled, a lighthouse can be started without a local tun interface (and therefore without root)
disabled: false
# Name of the device. If not set, a default will be chosen by the OS.
# For Linux: a single `%d` anywhere in the name is treated as a template and replaced with the
# lowest number that yields an unused device name (e.g. `nebula%d` becomes `nebula0`, then `nebula1`, and so on, `neb%dprod` becomes `neb0prod`).
# Only on Linux: `nebula%d` is the default if tun.dev is unset.
# The name, both before and after %d substitution, must be shorter than the kernel limit of 16 characters.
# For macOS: if set, must be in the form `utun[0-9]+`.
# For NetBSD: Required to be set, must be in the form `tun[0-9]+`
dev: nebula1
@@ -265,25 +258,6 @@ tun:
# Default MTU for every packet, safe setting is (and the default) 1300 for internet based traffic
mtu: 1300
# Linux only. pin_threads pins each tun reader/encrypt OS thread to a single CPU. This keeps every goroutine's
# batched sends flowing through one XPS-selected NIC TX ring, so packets within a flow stay ordered on the wire
# instead of being sprayed across multiple TX rings and reordered. Not reloadable. Coerced to false if routines <= 1.
#pin_threads: true
# Linux only. cpu_affinity overrides which CPUs the tun reader threads pin to: a list of CPU IDs, one per routine
# (see the top-level `routines` setting). Lists shorter than `routines` are modulo-cycled across the queues; extra
# entries are ignored. IDs must be within the process's allowed CPU set, so this respects taskset / cgroup cpusets;
# a non-integer or not-allowed entry disables the override and falls back to spreading queues across the allowed
# CPUs. Only meaningful while pin_threads is true. Not reloadable.
# When unset, the default spread prefers performance cores on heterogeneous CPUs (ARM big.LITTLE, Intel P/E
# hybrids, AMD compact cores), keeps all readers on one NUMA node and on distinct physical cores when the
# topology allows (SMT siblings last), leaves CPU 0's physical core as a last resort, and rotates its starting
# point per instance (keyed by the bound UDP port) so co-located nebulas don't stack their readers onto the
# same cores.
#cpu_affinity:
# - 2
# - 4
# Route based MTU overrides, you have known vpn ip paths that can support larger MTUs you can increase/decrease them here
routes:
#- mtu: 8800
-9
View File
@@ -8,15 +8,6 @@ Before=sshd.service
Type=notify
NotifyAccess=main
SyslogIdentifier=nebula
# Uncomment to run as an unprivileged user with only CAP_NET_ADMIN. Requires a
# nebula user that owns the config directory. Add CAP_NET_BIND_SERVICE to both
# lines if any listener (lighthouse DNS, listen.port, stats, sshd) binds <1024.
#User=nebula
#Group=nebula
#CapabilityBoundingSet=CAP_NET_ADMIN
#AmbientCapabilities=CAP_NET_ADMIN
ExecReload=/bin/kill -HUP $MAINPID
ExecStart=/usr/local/bin/nebula -config /etc/nebula/config.yml
Restart=always
+8 -8
View File
@@ -44,8 +44,8 @@ type Firewall struct {
InRules *FirewallTable
OutRules *FirewallTable
InboundSendReject bool
OutboundSendReject bool
InSendReject bool
OutSendReject bool
//TODO: we should have many more options for TCP, an option for ICMP, and mimic the kernel a bit better
// https://www.kernel.org/doc/Documentation/networking/nf_conntrack-sysctl.txt
@@ -216,23 +216,23 @@ func NewFirewallFromConfig(l *slog.Logger, cs *CertState, c *config.C) (*Firewal
inboundAction := c.GetString("firewall.inbound_action", "drop")
switch inboundAction {
case "reject":
fw.InboundSendReject = true
fw.InSendReject = true
case "drop":
fw.InboundSendReject = false
fw.InSendReject = false
default:
l.Warn("invalid firewall.inbound_action, defaulting to `drop`", "action", inboundAction)
fw.InboundSendReject = false
fw.InSendReject = false
}
outboundAction := c.GetString("firewall.outbound_action", "drop")
switch outboundAction {
case "reject":
fw.OutboundSendReject = true
fw.OutSendReject = true
case "drop":
fw.OutboundSendReject = false
fw.OutSendReject = false
default:
l.Warn("invalid firewall.outbound_action, defaulting to `drop`", "action", outboundAction)
fw.OutboundSendReject = false
fw.OutSendReject = false
}
err := AddFirewallRulesFromConfig(l, false, c, fw)
+2 -4
View File
@@ -5,8 +5,6 @@ import (
"log/slog"
"sync/atomic"
"time"
"github.com/slackhq/nebula/logging"
)
// ConntrackCache is used as a local routine cache to know if a given flow
@@ -58,8 +56,8 @@ func (c *ConntrackCacheTicker) Get() ConntrackCache {
if tick := c.cacheTick.Load(); tick != c.cacheV {
c.cacheV = tick
if ll := len(c.cache); ll > 0 {
if c.l.Enabled(context.Background(), logging.LevelTrace) {
c.l.Log(context.Background(), logging.LevelTrace, "resetting conntrack cache", "len", ll)
if c.l.Enabled(context.Background(), slog.LevelDebug) {
c.l.Debug("resetting conntrack cache", "len", ll)
}
c.cache = make(ConntrackCache, ll)
}
+7 -8
View File
@@ -6,7 +6,6 @@ import (
"strings"
"testing"
"github.com/slackhq/nebula/logging"
"github.com/slackhq/nebula/test"
"github.com/stretchr/testify/assert"
)
@@ -31,27 +30,27 @@ func newFixedTicker(t *testing.T, l *slog.Logger, cacheLen int) *ConntrackCacheT
func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) {
buf := &bytes.Buffer{}
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
c := newFixedTicker(t, l, 3)
c.Get()
assert.Equal(t, "level=DEBUG-4 msg=\"resetting conntrack cache\" len=3\n", buf.String())
assert.Equal(t, "level=DEBUG msg=\"resetting conntrack cache\" len=3\n", buf.String())
}
func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) {
buf := &bytes.Buffer{}
l := test.NewJSONLoggerWithOutput(buf, logging.LevelTrace)
l := test.NewJSONLoggerWithOutput(buf, slog.LevelDebug)
c := newFixedTicker(t, l, 2)
c.Get()
assert.JSONEq(t, `{"level":"DEBUG-4","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
assert.JSONEq(t, `{"level":"DEBUG","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
}
func TestConntrackCacheTicker_Get_QuietBelowTrace(t *testing.T) {
func TestConntrackCacheTicker_Get_QuietBelowDebug(t *testing.T) {
buf := &bytes.Buffer{}
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelInfo)
c := newFixedTicker(t, l, 5)
c.Get()
@@ -61,7 +60,7 @@ func TestConntrackCacheTicker_Get_QuietBelowTrace(t *testing.T) {
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
buf := &bytes.Buffer{}
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
c := newFixedTicker(t, l, 0)
c.Get()
-9
View File
@@ -65,12 +65,3 @@ func (fp Packet) MarshalJSON() ([]byte, error) {
"Fragment": fp.Fragment,
})
}
// ParsedPacket is a Packet plus the parse byproducts the RX path reuses
type ParsedPacket struct {
Packet
IPHdrLen int
// FragAny reports any fragmentation at all: MF flag or nonzero offset for IPv4, a fragment extension header for IPv6.
// Distinct from Packet.Fragment, which is true only for NON-FIRST fragments
FragAny bool
}
+6 -6
View File
@@ -1,6 +1,6 @@
module github.com/slackhq/nebula
go 1.26.0
go 1.25.0
require (
dario.cat/mergo v1.0.2
@@ -24,12 +24,12 @@ require (
github.com/vishvananda/netlink v1.3.1
go.uber.org/goleak v1.3.0
go.yaml.in/yaml/v3 v3.0.4
golang.org/x/crypto v0.54.0
golang.org/x/crypto v0.53.0
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
golang.org/x/net v0.57.0
golang.org/x/sync v0.22.0
golang.org/x/sys v0.47.0
golang.org/x/term v0.45.0
golang.org/x/net v0.56.0
golang.org/x/sync v0.21.0
golang.org/x/sys v0.46.0
golang.org/x/term v0.44.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b
golang.zx2c4.com/wireguard/windows v1.0.1
+10 -10
View File
@@ -162,8 +162,8 @@ golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACk
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw=
golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk=
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
@@ -182,8 +182,8 @@ golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLL
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE=
golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU=
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
@@ -191,8 +191,8 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM=
golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
@@ -208,11 +208,11 @@ golang.org/x/sys v0.0.0-20210124154548-22da62e12c0c/go.mod h1:h1NjWce9XRLGQEsW7w
golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc=
golang.org/x/term v0.44.0/go.mod h1:7ze4MdzUzLXpSAoFP1H0bOI9aXDqveSvatT5vKcFh2Y=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
+9 -17
View File
@@ -295,13 +295,7 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
hm.messageMetrics.Tx(header.Handshake, hh.machine.Subtype(), 1)
err := hm.outside.WriteTo(stage0, addr)
if err != nil {
// These repeat every attempt, so match the success log below and only shout when the remotes changed
level := slog.LevelDebug
if remotesHaveChanged {
level = slog.LevelError
}
hostinfo.logger(hm.l).Log(context.Background(), level, "Failed to send handshake message",
hostinfo.logger(hm.l).Error("Failed to send handshake message",
"udpAddr", addr,
"initiatorIndex", hostinfo.localIndexId,
"handshake", hsFields,
@@ -436,11 +430,14 @@ func (hm *HandshakeManager) CheckAndComplete(hostinfo *HostInfo, handshakePacket
// Check if we already have a tunnel with this vpn ip
existingHostInfo, found := hm.mainHostMap.Hosts[hostinfo.vpnAddrs[0]]
if found && existingHostInfo != nil {
// Is it just a delayed handshake packet? Check every hostinfo we hold for this address.
for _, testHostInfo := range hm.mainHostMap.unlockedGetHostList(hostinfo.vpnAddrs[0]) {
testHostInfo := existingHostInfo
for testHostInfo != nil {
// Is it just a delayed handshake packet?
if bytes.Equal(hostinfo.HandshakePacket[handshakePacket], testHostInfo.HandshakePacket[handshakePacket]) {
return testHostInfo, ErrAlreadySeen
}
testHostInfo = testHostInfo.next
}
// Is this a newer handshake?
@@ -535,9 +532,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 +970,6 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
nb := make([]byte, 12, 12)
out := make([]byte, mtu)
for _, cp := range hh.packetStore {
// TODO: use a SendBatch here. Each callback lands in
// sendNoMetrics -> WriteTo: one syscall per cached packet,
// where one sendmmsg could flush the whole store.
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
}
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
@@ -1088,8 +1080,8 @@ func (hm *HandshakeManager) sendHandshakeResponse(via ViaSender, msg []byte, hos
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
// We received a valid handshake on this relay, so make sure the relay
// state reflects that, in case it had been marked Disestablished.
via.relayHI.relayState.UpdateRelayForByIdxState(via.relay.LocalIndex, Established)
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false, 0)
via.relayHI.relayState.UpdateRelayForByIdxState(via.remoteIdx, Established)
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...)
}
}
+1 -1
View File
@@ -84,7 +84,7 @@ func (mw *mockEncWriter) SendMessageToVpnAddr(_ header.MessageType, _ header.Mes
return
}
func (mw *mockEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
func (mw *mockEncWriter) SendVia(_ *HostInfo, _ *Relay, _, _, _ []byte, _ bool) {
return
}
+6 -11
View File
@@ -190,18 +190,13 @@ func SubTypeName(t MessageType, s MessageSubType) string {
}
func IsValidSubType(t MessageType, s MessageSubType) bool {
switch t {
case Message:
return s == MessageNone || s == MessageRelay
case Handshake:
return s == HandshakeIXPSK0
case Test:
return s == TestReply || s == TestRequest
case Control, CloseTunnel, RecvError, LightHouse:
return s == 0
default:
return false
if n, ok := subTypeMap[t]; ok {
if _, ok := (*n)[s]; ok {
return true
}
}
return false
}
// NewHeader turns bytes into a header
-51
View File
@@ -102,57 +102,6 @@ func TestTypeMap(t *testing.T) {
}, subTypeMap)
}
// mapIsValidSubType is the pre-refactor, map-driven definition of a valid
// subtype. IsValidSubType was reimplemented as an explicit switch; this keeps
// the original behavior around so we can prove the switch is equivalent to it.
func mapIsValidSubType(t MessageType, s MessageSubType) bool {
if n, ok := subTypeMap[t]; ok {
if _, ok := (*n)[s]; ok {
return true
}
}
return false
}
func TestIsValidSubType(t *testing.T) {
// Explicit intent table: documents exactly which subtypes are valid so the
// test stays meaningful even if both the switch and subTypeMap change.
assert.True(t, IsValidSubType(Message, MessageNone))
assert.True(t, IsValidSubType(Message, MessageRelay))
assert.False(t, IsValidSubType(Message, 2))
assert.True(t, IsValidSubType(Handshake, HandshakeIXPSK0))
// HandshakeXXPSK0 is defined but not a wire-valid subtype.
assert.False(t, IsValidSubType(Handshake, HandshakeXXPSK0))
assert.True(t, IsValidSubType(Test, TestRequest))
assert.True(t, IsValidSubType(Test, TestReply))
assert.False(t, IsValidSubType(Test, 2))
// These types only ever carry subtype 0.
for _, mt := range []MessageType{Control, CloseTunnel, RecvError, LightHouse} {
assert.True(t, IsValidSubType(mt, 0), "type %d subtype 0 should be valid", mt)
assert.False(t, IsValidSubType(mt, 1), "type %d subtype 1 should be invalid", mt)
}
// Unknown/unassigned types are never valid.
assert.False(t, IsValidSubType(99, 0))
// Exhaustive proof of equivalence with the original map-driven logic across
// the entire (type, subtype) input space.
for ti := 0; ti <= 0xff; ti++ {
for si := 0; si <= 0xff; si++ {
mt, mst := MessageType(ti), MessageSubType(si)
assert.Equalf(t, mapIsValidSubType(mt, mst), IsValidSubType(mt, mst),
"IsValidSubType(%d, %d) diverged from map-driven definition", ti, si)
}
}
// H method must delegate to the package function.
assert.True(t, (&H{Type: Test, Subtype: TestReply}).IsValidSubType())
assert.False(t, (&H{Type: Handshake, Subtype: HandshakeXXPSK0}).IsValidSubType())
}
func TestHeader_String(t *testing.T) {
assert.Equal(
t,
+97 -162
View File
@@ -56,20 +56,11 @@ type Relay struct {
}
type HostMap struct {
sync.RWMutex //Because we concurrently read and write to our maps
Indexes map[uint32]*HostInfo
Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object
RemoteIndexes map[uint32]*HostInfo
// Hosts maps a vpn address to its primary hostinfo, one entry per address we hold a tunnel
// for. moreHosts only has an entry while an address is held by 2 or more hostinfos and stores
// the full most-recent-first list; moreHosts[a][0] is always the same hostinfo as Hosts[a].
// Each address gets its own independent list, so a hostinfo owning multiple addresses can
// never corrupt another address's ordering the way the old shared next/prev chain could.
// Entries in moreHosts are only ever written by unlockedSetHostsForAddr; Hosts is written
// directly only in the single-hostinfo fast paths where moreHosts is known to have no entry,
// and unlockedDeleteHostInfo swaps either map for a fresh one when it fully drains.
sync.RWMutex //Because we concurrently read and write to our maps
Indexes map[uint32]*HostInfo
Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object
RemoteIndexes map[uint32]*HostInfo
Hosts map[netip.Addr]*HostInfo
moreHosts map[netip.Addr][]*HostInfo
preferredRanges atomic.Pointer[[]netip.Prefix]
l *slog.Logger
}
@@ -275,6 +266,10 @@ type HostInfo struct {
lastRoam time.Time
lastRoamRemote netip.AddrPort
// Used to track other hostinfos for this vpn ip since only 1 can be primary
// Synchronised via hostmap lock and not the hostinfo lock.
next, prev *HostInfo
//TODO: in, out, and others might benefit from being an atomic.Int32. We could collapse connectionManager pendingDeletion, relayUsed, and in/out into this 1 thing
in, out, pendingDeletion atomic.Bool
@@ -287,6 +282,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
}
@@ -338,7 +334,6 @@ func newHostMap(l *slog.Logger) *HostMap {
Relays: map[uint32]*HostInfo{},
RemoteIndexes: map[uint32]*HostInfo{},
Hosts: map[netip.Addr]*HostInfo{},
moreHosts: map[netip.Addr][]*HostInfo{},
l: l,
}
}
@@ -387,55 +382,13 @@ func (hm *HostMap) EmitStats() {
metrics.GetOrRegisterGauge("hostmap.main.relayIndexes", nil).Update(int64(relaysLen))
}
// unlockedSetHostsForAddr stores the per-address hostinfo list (list[0] is the primary). An empty
// list removes the address. This is the one place Hosts and moreHosts are written together, keep
// it that way. Callers must hold the write lock.
func (hm *HostMap) unlockedSetHostsForAddr(addr netip.Addr, list []*HostInfo) {
if len(list) == 0 {
delete(hm.Hosts, addr)
delete(hm.moreHosts, addr)
return
}
hm.Hosts[addr] = list[0]
if len(list) > 1 {
hm.moreHosts[addr] = list
} else {
delete(hm.moreHosts, addr)
}
}
// unlockedGetHostList returns every hostinfo holding addr, primary first, or nil if we have no
// tunnel for addr. The common single-hostinfo case builds a fresh one element list, so keep this
// off the packet hot path; the primary is a direct Hosts read. Callers must hold the lock (read
// or write).
func (hm *HostMap) unlockedGetHostList(addr netip.Addr) []*HostInfo {
if list, ok := hm.moreHosts[addr]; ok {
return list
}
if h, ok := hm.Hosts[addr]; ok {
return []*HostInfo{h}
}
return nil
}
// removeHostInfo returns list with hi removed (order preserved), or list unchanged if hi is
// absent. It deletes in place: every mutator holds the hostmap write lock and no reader ever
// retains a slice across a mutation (readers iterate under RLock), so there is no snapshot to
// invalidate.
func removeHostInfo(list []*HostInfo, hi *HostInfo) []*HostInfo {
idx := slices.Index(list, hi)
if idx < 0 {
return list
}
return slices.Delete(list, idx, idx+1)
}
// DeleteHostInfo will fully unlink the hostinfo and return true if no other hostinfo still holds
// any of its vpn addrs, meaning we no longer have a tunnel to the peer
// DeleteHostInfo will fully unlink the hostinfo and return true if it was the final hostinfo for this vpn ip
func (hm *HostMap) DeleteHostInfo(hostinfo *HostInfo) bool {
// Delete the host itself, ensuring it's not modified anymore
hm.Lock()
final := hm.unlockedDeleteHostInfo(hostinfo)
// If we have a previous or next hostinfo then we are not the last one for this vpn ip
final := (hostinfo.next == nil && hostinfo.prev == nil)
hm.unlockedDeleteHostInfo(hostinfo)
hm.Unlock()
return final
@@ -447,66 +400,71 @@ func (hm *HostMap) MakePrimary(hostinfo *HostInfo) {
hm.unlockedMakePrimary(hostinfo)
}
// unlockedMakePrimary reports whether hostinfo is (now) the primary for each of its addresses,
// false only when it is no longer in the hostmap at all.
func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) bool {
// A hostinfo that is no longer in the hostmap must not be re-inserted here. Callers can race
// tunnel teardown, deciding to promote under the read lock and only taking the write lock
// after a delete fully unlinked the hostinfo (connection manager swapPrimary, AddRelay). Every
// live hostinfo is registered in Indexes by unlockedAddHostInfo, so this is a membership test.
if hm.Indexes[hostinfo.localIndexId] != hostinfo {
return false
func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) {
// Get the current primary, if it exists
oldHostinfo := hm.Hosts[hostinfo.vpnAddrs[0]]
// Every address in the hostinfo gets elevated to primary
for _, vpnAddr := range hostinfo.vpnAddrs {
//NOTE: It is possible that we leave a dangling hostinfo here but connection manager works on
// indexes so it should be fine.
hm.Hosts[vpnAddr] = hostinfo
}
// Move hostinfo to the front (primary) of each of its address lists. The lists are
// independent per address, so this can never leave a dangling entry the way promoting
// against a single shared chain could.
for _, addr := range hostinfo.vpnAddrs {
if hm.Hosts[addr] == hostinfo {
// Already primary for this address, the list is already in the right order
continue
}
list := removeHostInfo(hm.unlockedGetHostList(addr), hostinfo)
list = append([]*HostInfo{hostinfo}, list...)
hm.unlockedSetHostsForAddr(addr, list)
// If we are already primary then we won't bother re-linking
if oldHostinfo == hostinfo {
return
}
return true
// Unlink this hostinfo
if hostinfo.prev != nil {
hostinfo.prev.next = hostinfo.next
}
if hostinfo.next != nil {
hostinfo.next.prev = hostinfo.prev
}
// If there wasn't a previous primary then clear out any links
if oldHostinfo == nil {
hostinfo.next = nil
hostinfo.prev = nil
return
}
// Relink the hostinfo as primary
hostinfo.next = oldHostinfo
oldHostinfo.prev = hostinfo
hostinfo.prev = nil
}
// unlockedDeleteHostInfo removes hostinfo from every one of its address lists and from the index
// maps. It returns true if this was the last hostinfo for all of its addresses (we no longer have
// any tunnel to the peer), which the caller uses to decide whether to clear learned lighthouse
// state and disestablish relays.
func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
// Remove this hostinfo from each of its address lists. The lists are independent, so a
// sibling is never promoted to an address it does not own and no other list is touched.
final := true
func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) {
isLastHostinfo := hostinfo.next == nil && hostinfo.prev == nil
for _, addr := range hostinfo.vpnAddrs {
if list, ok := hm.moreHosts[addr]; ok {
list = removeHostInfo(list, hostinfo)
hm.unlockedSetHostsForAddr(addr, list)
if len(list) > 0 {
final = false
}
} else if existing, ok := hm.Hosts[addr]; ok {
if existing == hostinfo {
// Common case, the only hostinfo for this address. moreHosts has no entry to clean up.
delete(hm.Hosts, addr)
} else {
// We don't hold this address but another hostinfo does, we still have a tunnel to the peer
final = false
}
if hm.Hosts[addr] != hostinfo {
continue
}
if hostinfo.next != nil {
// Promote the next hostinfo in the shared chain to primary for this address
hm.Hosts[addr] = hostinfo.next
} else {
delete(hm.Hosts, addr)
}
}
// Go maps never shrink their buckets, replace fully drained maps so a node that churned
// through a large peer count gives the memory back. Same idiom as the index maps below.
if len(hm.Hosts) == 0 {
hm.Hosts = map[netip.Addr]*HostInfo{}
}
if len(hm.moreHosts) == 0 {
hm.moreHosts = map[netip.Addr][]*HostInfo{}
// Splice this hostinfo out of the shared chain exactly once
if hostinfo.prev != nil {
hostinfo.prev.next = hostinfo.next
}
if hostinfo.next != nil {
hostinfo.next.prev = hostinfo.prev
}
hostinfo.next = nil
hostinfo.prev = nil
// The remote index uses index ids outside our control so lets make sure we are only removing
// the remote index pointer here if it points to the hostinfo we are deleting
@@ -530,7 +488,7 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
)
}
if final {
if isLastHostinfo {
// I have lost connectivity to my peers. My relay tunnel is likely broken. Mark the next
// hops as 'Requested' so that new relay tunnels are created in the future.
hm.unlockedDisestablishVpnAddrRelayFor(hostinfo)
@@ -539,19 +497,6 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
for _, localRelayIdx := range hostinfo.relayState.CopyRelayForIdxs() {
delete(hm.Relays, localRelayIdx)
}
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 {
@@ -595,30 +540,19 @@ func (hm *HostMap) QueryVpnAddrsRelayFor(targetIps []netip.Addr, relayHostIp net
hm.RLock()
defer hm.RUnlock()
// This runs per relayed packet, so check the primary with a single map probe and only consult
// moreHosts when the primary can't relay for us.
h, ok := hm.Hosts[relayHostIp]
if !ok {
return nil, nil, errors.New("unable to find host")
}
for _, targetIp := range targetIps {
r, ok := h.relayState.QueryRelayForByIp(targetIp)
if ok && r.State == Established {
return h, r, nil
}
}
if list, ok := hm.moreHosts[relayHostIp]; ok {
// list[0] is the primary we already checked
for _, h := range list[1:] {
for _, targetIp := range targetIps {
r, ok := h.relayState.QueryRelayForByIp(targetIp)
if ok && r.State == Established {
return h, r, nil
}
for h != nil {
for _, targetIp := range targetIps {
r, ok := h.relayState.QueryRelayForByIp(targetIp)
if ok && r.State == Established {
return h, r, nil
}
}
h = h.next
}
return nil, nil, errors.New("unable to find host with relay")
@@ -626,14 +560,20 @@ func (hm *HostMap) QueryVpnAddrsRelayFor(targetIps []netip.Addr, relayHostIp net
func (hm *HostMap) unlockedDisestablishVpnAddrRelayFor(hi *HostInfo) {
for _, relayHostIp := range hi.relayState.CopyRelayIps() {
for _, h := range hm.unlockedGetHostList(relayHostIp) {
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
if h, ok := hm.Hosts[relayHostIp]; ok {
for h != nil {
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
h = h.next
}
}
}
for _, rs := range hi.relayState.CopyAllRelayFor() {
if rs.Type == ForwardingType {
for _, h := range hm.unlockedGetHostList(rs.PeerAddr) {
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
if h, ok := hm.Hosts[rs.PeerAddr]; ok {
for h != nil {
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
h = h.next
}
}
}
}
@@ -683,27 +623,22 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
}
func (hm *HostMap) unlockedInnerAddHostInfo(vpnAddr netip.Addr, hostinfo *HostInfo, f *Interface) {
existing, ok := hm.Hosts[vpnAddr]
if !ok {
// Common case, the first hostinfo for this address. moreHosts stays empty.
hm.Hosts[vpnAddr] = hostinfo
return
existing := hm.Hosts[vpnAddr]
hm.Hosts[vpnAddr] = hostinfo
if existing != nil && existing != hostinfo {
hostinfo.next = existing
existing.prev = hostinfo
}
// The new hostinfo becomes the primary for this address. Remove any stale copy of it first so
// we never hold a duplicate, then prepend.
list, ok := hm.moreHosts[vpnAddr]
if !ok {
list = []*HostInfo{existing}
}
list = removeHostInfo(list, hostinfo)
list = append([]*HostInfo{hostinfo}, list...)
hm.unlockedSetHostsForAddr(vpnAddr, list)
// Enforce the per-address cap by fully retiring the oldest hostinfo once we exceed it.
// Deleting it removes it from all of its addresses and the index maps, matching prior behavior.
if len(list) > MaxHostInfosPerVpnIp {
hm.unlockedDeleteHostInfo(list[len(list)-1])
i := 1
check := hostinfo
for check != nil {
if i > MaxHostInfosPerVpnIp {
hm.unlockedDeleteHostInfo(check)
}
check = check.next
i++
}
}
+182 -238
View File
@@ -2,7 +2,6 @@ package nebula
import (
"net/netip"
"slices"
"testing"
"github.com/slackhq/nebula/config"
@@ -11,84 +10,78 @@ import (
"github.com/stretchr/testify/require"
)
// chainIds returns the localIndexIds of the hostinfos holding addr, primary (index 0) first. It
// also validates the Hosts/moreHosts sync contract on every call so a mutation that broke it
// fails fast.
func chainIds(t *testing.T, hm *HostMap, addr netip.Addr) []uint32 {
t.Helper()
assertHostMapInvariants(t, hm)
list := hm.unlockedGetHostList(addr)
ids := make([]uint32, len(list))
for i, h := range list {
ids[i] = h.localIndexId
}
return ids
}
// assertHostMapInvariants checks the Hosts/moreHosts contract: moreHosts only holds addresses
// with 2 or more hostinfos, its first entry is always the primary in Hosts, lists never hold
// duplicates, every hostinfo in a list owns the address and is registered in Indexes, and every
// indexed hostinfo is reachable through each of its addresses.
func assertHostMapInvariants(t *testing.T, hm *HostMap) {
t.Helper()
for addr, list := range hm.moreHosts {
require.GreaterOrEqualf(t, len(list), 2, "moreHosts[%s] must hold at least 2 hostinfos", addr)
require.Samef(t, hm.Hosts[addr], list[0], "moreHosts[%s][0] must match the primary in Hosts", addr)
seen := map[*HostInfo]bool{}
for _, h := range list {
require.NotNilf(t, h, "moreHosts[%s] must never hold a nil hostinfo", addr)
require.Falsef(t, seen[h], "moreHosts[%s] holds hostinfo %d twice", addr, h.localIndexId)
seen[h] = true
require.Samef(t, hm.Indexes[h.localIndexId], h, "moreHosts[%s] member %d is not registered in Indexes", addr, h.localIndexId)
require.Truef(t, slices.Contains(h.vpnAddrs, addr), "moreHosts[%s] member %d does not own the address", addr, h.localIndexId)
}
}
for addr, h := range hm.Hosts {
require.NotNilf(t, h, "Hosts[%s] must never be nil", addr)
require.Samef(t, hm.Indexes[h.localIndexId], h, "Hosts[%s] primary %d is not registered in Indexes", addr, h.localIndexId)
require.Truef(t, slices.Contains(h.vpnAddrs, addr), "Hosts[%s] primary (index %d) does not own the address", addr, h.localIndexId)
}
for idx, h := range hm.Indexes {
require.Equalf(t, idx, h.localIndexId, "Indexes[%d] holds hostinfo with localIndexId %d", idx, h.localIndexId)
for _, va := range h.vpnAddrs {
require.Truef(t, slices.Contains(hm.unlockedGetHostList(va), h), "indexed hostinfo %d is missing from the list for %s", idx, va)
}
}
}
func TestHostMap_MakePrimary(t *testing.T) {
l := test.NewLogger()
hm := newHostMap(l)
f := &Interface{}
a := netip.MustParseAddr("0.0.0.1")
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2}
h3 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 3}
h4 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 4}
h1 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 1}
h2 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 2}
h3 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 3}
h4 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 4}
hm.unlockedAddHostInfo(h4, f)
hm.unlockedAddHostInfo(h3, f)
hm.unlockedAddHostInfo(h2, f)
hm.unlockedAddHostInfo(h1, f)
// Most-recently-added is primary: h1, h2, h3, h4
assert.Equal(t, []uint32{1, 2, 3, 4}, chainIds(t, hm, a))
assert.Equal(t, h1, hm.QueryVpnAddr(a))
// Make sure we go h1 -> h2 -> h3 -> h4
prim := hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, h1.localIndexId, prim.localIndexId)
assert.Equal(t, h2.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
assert.Equal(t, h3.localIndexId, h2.next.localIndexId)
assert.Equal(t, h2.localIndexId, h3.prev.localIndexId)
assert.Equal(t, h4.localIndexId, h3.next.localIndexId)
assert.Equal(t, h3.localIndexId, h4.prev.localIndexId)
assert.Nil(t, h4.next)
// Swap the middle to primary: h3, h1, h2, h4
// Swap h3/middle to primary
hm.MakePrimary(h3)
assert.Equal(t, []uint32{3, 1, 2, 4}, chainIds(t, hm, a))
assert.Equal(t, h3, hm.QueryVpnAddr(a))
// Swap the tail to primary: h4, h3, h1, h2
hm.MakePrimary(h4)
assert.Equal(t, []uint32{4, 3, 1, 2}, chainIds(t, hm, a))
// Make sure we go h3 -> h1 -> h2 -> h4
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, h3.localIndexId, prim.localIndexId)
assert.Equal(t, h1.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h2.localIndexId, h1.next.localIndexId)
assert.Equal(t, h3.localIndexId, h1.prev.localIndexId)
assert.Equal(t, h4.localIndexId, h2.next.localIndexId)
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
assert.Equal(t, h2.localIndexId, h4.prev.localIndexId)
assert.Nil(t, h4.next)
// Swapping the current primary again is a no-op
// Swap h4/tail to primary
hm.MakePrimary(h4)
assert.Equal(t, []uint32{4, 3, 1, 2}, chainIds(t, hm, a))
// Make sure we go h4 -> h3 -> h1 -> h2
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, h4.localIndexId, prim.localIndexId)
assert.Equal(t, h3.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h1.localIndexId, h3.next.localIndexId)
assert.Equal(t, h4.localIndexId, h3.prev.localIndexId)
assert.Equal(t, h2.localIndexId, h1.next.localIndexId)
assert.Equal(t, h3.localIndexId, h1.prev.localIndexId)
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
assert.Nil(t, h2.next)
// Swap h4 again should be no-op
hm.MakePrimary(h4)
// Make sure we go h4 -> h3 -> h1 -> h2
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, h4.localIndexId, prim.localIndexId)
assert.Equal(t, h3.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h1.localIndexId, h3.next.localIndexId)
assert.Equal(t, h4.localIndexId, h3.prev.localIndexId)
assert.Equal(t, h2.localIndexId, h1.next.localIndexId)
assert.Equal(t, h3.localIndexId, h1.prev.localIndexId)
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
assert.Nil(t, h2.next)
}
func TestHostMap_DeleteHostInfo(t *testing.T) {
@@ -96,14 +89,13 @@ func TestHostMap_DeleteHostInfo(t *testing.T) {
hm := newHostMap(l)
f := &Interface{}
a := netip.MustParseAddr("0.0.0.1")
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2}
h3 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 3}
h4 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 4}
h5 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 5}
h6 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 6}
h1 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 1}
h2 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 2}
h3 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 3}
h4 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 4}
h5 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 5}
h6 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 6}
hm.unlockedAddHostInfo(h6, f)
hm.unlockedAddHostInfo(h5, f)
@@ -112,110 +104,94 @@ func TestHostMap_DeleteHostInfo(t *testing.T) {
hm.unlockedAddHostInfo(h2, f)
hm.unlockedAddHostInfo(h1, f)
// h6 is evicted by the MaxHostInfosPerVpnIp cap; the rest are newest-first.
assert.Nil(t, hm.QueryIndex(h6.localIndexId))
assert.Equal(t, []uint32{1, 2, 3, 4, 5}, chainIds(t, hm, a))
// h6 should be deleted
assert.Nil(t, h6.next)
assert.Nil(t, h6.prev)
h := hm.QueryIndex(h6.localIndexId)
assert.Nil(t, h)
// Delete primary; not final since siblings remain.
assert.False(t, hm.DeleteHostInfo(h1))
assert.Equal(t, []uint32{2, 3, 4, 5}, chainIds(t, hm, a))
// Make sure we go h1 -> h2 -> h3 -> h4 -> h5
prim := hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, h1.localIndexId, prim.localIndexId)
assert.Equal(t, h2.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
assert.Equal(t, h3.localIndexId, h2.next.localIndexId)
assert.Equal(t, h2.localIndexId, h3.prev.localIndexId)
assert.Equal(t, h4.localIndexId, h3.next.localIndexId)
assert.Equal(t, h3.localIndexId, h4.prev.localIndexId)
assert.Equal(t, h5.localIndexId, h4.next.localIndexId)
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
assert.Nil(t, h5.next)
// Deleting the same hostinfo again must not report final while siblings remain and must not
// disturb the list. The old chain code got this wrong: the first delete nil'd next/prev, so a
// second delete looked final and wiped lighthouse state out from under the live sibling.
assert.False(t, hm.DeleteHostInfo(h1))
assert.Equal(t, []uint32{2, 3, 4, 5}, chainIds(t, hm, a))
// Delete primary
hm.DeleteHostInfo(h1)
assert.Nil(t, h1.prev)
assert.Nil(t, h1.next)
// Delete a middle node.
assert.False(t, hm.DeleteHostInfo(h3))
assert.Equal(t, []uint32{2, 4, 5}, chainIds(t, hm, a))
// Make sure we go h2 -> h3 -> h4 -> h5
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, h2.localIndexId, prim.localIndexId)
assert.Equal(t, h3.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h3.localIndexId, h2.next.localIndexId)
assert.Equal(t, h2.localIndexId, h3.prev.localIndexId)
assert.Equal(t, h4.localIndexId, h3.next.localIndexId)
assert.Equal(t, h3.localIndexId, h4.prev.localIndexId)
assert.Equal(t, h5.localIndexId, h4.next.localIndexId)
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
assert.Nil(t, h5.next)
// Delete the tail.
assert.False(t, hm.DeleteHostInfo(h5))
assert.Equal(t, []uint32{2, 4}, chainIds(t, hm, a))
// Delete in the middle
hm.DeleteHostInfo(h3)
assert.Nil(t, h3.prev)
assert.Nil(t, h3.next)
// Delete the head; h4 remains and becomes primary.
assert.False(t, hm.DeleteHostInfo(h2))
assert.Equal(t, []uint32{4}, chainIds(t, hm, a))
assert.Equal(t, h4, hm.QueryVpnAddr(a))
// Make sure we go h2 -> h4 -> h5
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, h2.localIndexId, prim.localIndexId)
assert.Equal(t, h4.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h4.localIndexId, h2.next.localIndexId)
assert.Equal(t, h2.localIndexId, h4.prev.localIndexId)
assert.Equal(t, h5.localIndexId, h4.next.localIndexId)
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
assert.Nil(t, h5.next)
// Delete the only remaining item; final is true and the address is gone.
assert.True(t, hm.DeleteHostInfo(h4))
assert.Empty(t, chainIds(t, hm, a))
assert.Nil(t, hm.QueryVpnAddr(a))
// Delete the tail
hm.DeleteHostInfo(h5)
assert.Nil(t, h5.prev)
assert.Nil(t, h5.next)
// Deleting an already-gone hostinfo is still final; nothing holds the address anymore.
assert.True(t, hm.DeleteHostInfo(h4))
assert.Empty(t, chainIds(t, hm, a))
}
// Make sure we go h2 -> h4
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, h2.localIndexId, prim.localIndexId)
assert.Equal(t, h4.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h4.localIndexId, h2.next.localIndexId)
assert.Equal(t, h2.localIndexId, h4.prev.localIndexId)
assert.Nil(t, h4.next)
// TestHostMap_MakePrimary_DeletedHostInfo covers promoting a hostinfo that lost a race with
// tunnel teardown: swapPrimary and AddRelay decide to promote while holding a stale pointer and
// only take the write lock after a delete fully unlinked the hostinfo. MakePrimary must be a
// no-op, not a resurrection that installs an unmanaged primary.
func TestHostMap_MakePrimary_DeletedHostInfo(t *testing.T) {
l := test.NewLogger()
hm := newHostMap(l)
f := &Interface{}
a := netip.MustParseAddr("0.0.0.1")
// Delete the head
hm.DeleteHostInfo(h2)
assert.Nil(t, h2.prev)
assert.Nil(t, h2.next)
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2}
hm.unlockedAddHostInfo(h1, f)
hm.unlockedAddHostInfo(h2, f)
// Make sure we only have h4
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, h4.localIndexId, prim.localIndexId)
assert.Nil(t, prim.prev)
assert.Nil(t, prim.next)
assert.Nil(t, h4.next)
// h1 is fully deleted while another goroutine still holds a pointer to it.
assert.False(t, hm.DeleteHostInfo(h1))
assert.Equal(t, []uint32{2}, chainIds(t, hm, a))
// Delete the only item
hm.DeleteHostInfo(h4)
assert.Nil(t, h4.prev)
assert.Nil(t, h4.next)
// The stale promote must not bring it back.
hm.MakePrimary(h1)
assert.Equal(t, []uint32{2}, chainIds(t, hm, a))
assert.Equal(t, h2, hm.QueryVpnAddr(a))
assert.Nil(t, hm.QueryIndex(h1.localIndexId))
}
// TestHostMap_QueryVpnAddrsRelayFor_NonPrimary makes sure a relay established on an older
// hostinfo is still found after a newer tunnel without relay state takes primary for the same
// address. The lookup checks the primary first and falls back to the rest of the list.
func TestHostMap_QueryVpnAddrsRelayFor_NonPrimary(t *testing.T) {
l := test.NewLogger()
hm := newHostMap(l)
f := &Interface{}
relayAddr := netip.MustParseAddr("0.0.0.9")
target := netip.MustParseAddr("0.0.0.1")
older := &HostInfo{
vpnAddrs: []netip.Addr{relayAddr},
localIndexId: 1,
relayState: RelayState{
relayForByAddr: map[netip.Addr]*Relay{},
relayForByIdx: map[uint32]*Relay{},
},
}
older.relayState.InsertRelay(target, 100, &Relay{Type: ForwardingType, State: Established, LocalIndex: 100, PeerAddr: target})
hm.unlockedAddHostInfo(older, f)
// The relay is found on the primary.
h, r, err := hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, relayAddr)
require.NoError(t, err)
assert.Equal(t, older, h)
assert.Equal(t, uint32(100), r.LocalIndex)
// A re-handshake with no relay state takes primary; the established relay on the older
// hostinfo must still be found through the fallback.
newer := &HostInfo{vpnAddrs: []netip.Addr{relayAddr}, localIndexId: 2}
hm.unlockedAddHostInfo(newer, f)
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, relayAddr))
h, r, err = hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, relayAddr)
require.NoError(t, err)
assert.Equal(t, older, h)
assert.Equal(t, uint32(100), r.LocalIndex)
// No hostinfo at all is a plain miss.
_, _, err = hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, netip.MustParseAddr("0.0.0.42"))
require.Error(t, err)
// Make sure we have nil
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Nil(t, prim)
}
// TestHostMap_DeleteHostInfo_MultipleVpnAddrs exercises the case where a hostinfo carries more than one
@@ -240,82 +216,32 @@ func TestHostMap_DeleteHostInfo_MultipleVpnAddrs(t *testing.T) {
hm.unlockedAddHostInfo(other, f)
hm.unlockedAddHostInfo(head, f)
// head is primary for both addresses, other is next in each address's list.
assert.Equal(t, head, hm.QueryVpnAddr(a))
assert.Equal(t, head, hm.QueryVpnAddr(b))
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, a))
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, b))
// head is primary for both addresses, other is next in the shared chain
assert.Equal(t, head.localIndexId, hm.QueryVpnAddr(a).localIndexId)
assert.Equal(t, head.localIndexId, hm.QueryVpnAddr(b).localIndexId)
assert.Equal(t, other.localIndexId, head.next.localIndexId)
assert.Equal(t, head.localIndexId, other.prev.localIndexId)
// Delete the head. other is still live, so it must become primary for BOTH addresses.
assert.False(t, hm.DeleteHostInfo(head))
assert.Equal(t, other, hm.QueryVpnAddr(a))
assert.Equal(t, other, hm.QueryVpnAddr(b))
assert.Equal(t, []uint32{1}, chainIds(t, hm, a))
assert.Equal(t, []uint32{1}, chainIds(t, hm, b))
hm.DeleteHostInfo(head)
// head is fully removed from the index map.
// Pre-fix: QueryVpnAddr(b) came back nil here because the second address was deleted rather than
// promoted, leaving other unreachable at b.
require.NotNil(t, hm.QueryVpnAddr(a))
require.NotNil(t, hm.QueryVpnAddr(b))
assert.Equal(t, other.localIndexId, hm.QueryVpnAddr(a).localIndexId)
assert.Equal(t, other.localIndexId, hm.QueryVpnAddr(b).localIndexId)
// other is now the only hostinfo in the chain
assert.Nil(t, other.prev)
assert.Nil(t, other.next)
// head is fully detached
assert.Nil(t, head.prev)
assert.Nil(t, head.next)
assert.Nil(t, hm.QueryIndex(head.localIndexId))
}
// TestHostMap_DeleteHostInfo_DivergentVpnAddrs covers chained hostinfos for the same peer whose
// vpnAddrs sets differ (a re-handshake cert added a second address). Deleting the superset node
// must not promote a sibling to an address it does not own.
func TestHostMap_DeleteHostInfo_DivergentVpnAddrs(t *testing.T) {
l := test.NewLogger()
hm := newHostMap(l)
f := &Interface{}
a := netip.MustParseAddr("0.0.0.1")
b := netip.MustParseAddr("0.0.0.2")
// sub owns only a; super (a newer handshake) owns a and b.
sub := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
super := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 2}
hm.unlockedAddHostInfo(sub, f)
hm.unlockedAddHostInfo(super, f)
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, a))
assert.Equal(t, []uint32{2}, chainIds(t, hm, b))
// Delete super: a promotes to sub (which owns it); b has no remaining owner and must be
// removed, not dangled at sub (which does not own b).
assert.False(t, hm.DeleteHostInfo(super))
assert.Equal(t, []uint32{1}, chainIds(t, hm, a))
assert.Empty(t, chainIds(t, hm, b))
assert.Equal(t, sub, hm.QueryVpnAddr(a))
assert.Nil(t, hm.QueryVpnAddr(b))
assert.Nil(t, hm.QueryIndex(super.localIndexId))
// Deleting sub cleans up fully.
assert.True(t, hm.DeleteHostInfo(sub))
assert.Nil(t, hm.QueryVpnAddr(a))
assertHostMapInvariants(t, hm)
}
// TestHostMap_AddDivergentOverlap covers a new hostinfo claiming addresses currently owned by two
// DIFFERENT hostinfos. The old single shared next/prev chain overwrote a pointer and orphaned one
// of them (in Indexes but unreachable via its address); independent per-address lists cannot.
func TestHostMap_AddDivergentOverlap(t *testing.T) {
l := test.NewLogger()
hm := newHostMap(l)
f := &Interface{}
a := netip.MustParseAddr("0.0.0.1")
b := netip.MustParseAddr("0.0.0.2")
hiA := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
hiP := &HostInfo{vpnAddrs: []netip.Addr{b}, localIndexId: 2}
hm.unlockedAddHostInfo(hiA, f)
hm.unlockedAddHostInfo(hiP, f)
hiB := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 3}
hm.unlockedAddHostInfo(hiB, f)
assert.Equal(t, []uint32{3, 1}, chainIds(t, hm, a))
assert.Equal(t, []uint32{3, 2}, chainIds(t, hm, b))
// hiA is still reachable via its address (not orphaned) and still indexed.
assert.Contains(t, chainIds(t, hm, a), hiA.localIndexId)
assert.NotNil(t, hm.QueryIndex(hiA.localIndexId))
}
// TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs verifies the MaxHostInfosPerVpnIp overflow prune
// (unlockedInnerAddHostInfo calls unlockedDeleteHostInfo on the oldest node once the chain is too long)
// still behaves when hostinfos carry more than one vpnAddr. The pruned node is always the tail, so it is
@@ -341,14 +267,32 @@ func TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs(t *testing.T) {
oldest := hostinfos[len(hostinfos)-1]
// The oldest hostinfo was pruned from both lists and the index map.
// The oldest hostinfo should have been pruned and fully detached
assert.Nil(t, oldest.next)
assert.Nil(t, oldest.prev)
assert.Nil(t, hm.QueryIndex(oldest.localIndexId))
// Both addresses hold exactly MaxHostInfosPerVpnIp survivors in the same order; oldest is absent.
require.Len(t, chainIds(t, hm, a), MaxHostInfosPerVpnIp)
assert.Equal(t, chainIds(t, hm, a), chainIds(t, hm, b), "both addresses must list the same survivors in the same order")
assert.NotContains(t, chainIds(t, hm, a), oldest.localIndexId)
assert.Equal(t, hm.QueryVpnAddr(a), hm.QueryVpnAddr(b))
// Both addresses resolve to the same head, and that head is one of the survivors (not the pruned one)
primA := hm.QueryVpnAddr(a)
primB := hm.QueryVpnAddr(b)
require.NotNil(t, primA)
require.NotNil(t, primB)
assert.Equal(t, primA.localIndexId, primB.localIndexId)
assert.NotEqual(t, oldest.localIndexId, primA.localIndexId)
// Walk the shared chain: exactly MaxHostInfosPerVpnIp survivors, no cycles, oldest absent
seen := map[uint32]struct{}{}
for h := primA; h != nil; h = h.next {
_, dup := seen[h.localIndexId]
require.False(t, dup, "cycle detected in hostinfo chain")
seen[h.localIndexId] = struct{}{}
if h.next != nil {
assert.Equal(t, h.localIndexId, h.next.prev.localIndexId, "prev pointer must mirror next")
}
}
assert.Len(t, seen, MaxHostInfosPerVpnIp)
_, prunedStillPresent := seen[oldest.localIndexId]
assert.False(t, prunedStillPresent)
}
func TestHostMap_reload(t *testing.T) {
+36 -201
View File
@@ -2,7 +2,6 @@ package nebula
import (
"context"
"io"
"log/slog"
"net/netip"
@@ -10,24 +9,10 @@ import (
"github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/iputil"
"github.com/slackhq/nebula/noiseutil"
"github.com/slackhq/nebula/overlay/batch"
"github.com/slackhq/nebula/overlay/tio"
"github.com/slackhq/nebula/routing"
)
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.ParsedPacket, nb []byte, sendBatch *batch.SendBatch, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
// borrowed: pkt.Bytes is owned by the originating tio.Queue and is
// only valid until the next Read on that queue. Every consumer below
// (parse, self-forward, handshake cache, sendInsideMessage) reads it
// synchronously; do not retain pkt outside this call. If a future
// caller needs to keep the packet, use pkt.Clone() to detach it from
// the borrow.
//
// pkt.Bytes is either one IP datagram (GSO zero) or a TSO/USO
// superpacket. In both cases the L3+L4 headers at the start describe
// the same 5-tuple every segment will share, so a single newPacket /
// firewall check covers the whole superpacket.
packet := pkt.Bytes
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) {
err := newPacket(packet, false, fwPacket)
if err != nil {
if f.l.Enabled(context.Background(), slog.LevelDebug) {
@@ -52,14 +37,7 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Parse
// routes packets from the Nebula addr to the Nebula addr through the Nebula
// TUN device.
if immediatelyForwardToSelf {
// Write copies into the kernel queue synchronously, so seg's lifetime ends at return.
// A self-forwarded superpacket would be re-handed to the
// kernel as one giant blob; segment first so the loopback
// path sees one IP datagram per Write.
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
_, werr := f.queues[q].Write(seg)
return werr
})
_, err := f.readers[q].Write(packet)
if err != nil {
f.l.Error("Failed to forward to tun", "error", err)
}
@@ -74,24 +52,12 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Parse
return
}
hostinfo, ready := f.getOrHandshakeConsiderRouting(&fwPacket.Packet, func(hh *HandshakeHostInfo) {
// borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt
// bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket),
// so retaining segments past the loop is safe.
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
hh.cachePacket(f.l, header.Message, 0, seg, f.sendMessageNow, f.cachedPacketMetrics)
return nil
})
if err != nil && f.l.Enabled(context.Background(), slog.LevelDebug) {
f.l.Debug("Failed to segment superpacket for handshake cache",
"error", err,
"vpnAddr", fwPacket.RemoteAddr,
)
}
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
hh.cachePacket(f.l, header.Message, 0, packet, f.sendMessageNow, f.cachedPacketMetrics)
})
if hostinfo == nil {
f.rejectInside(packet, rejectBuf, q)
f.rejectInside(packet, out, q)
if f.l.Enabled(context.Background(), slog.LevelDebug) {
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
"vpnAddr", fwPacket.RemoteAddr,
@@ -105,11 +71,12 @@ 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.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q)
} else {
f.rejectInside(packet, rejectBuf, q)
f.rejectInside(packet, out, q)
if f.l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(f.l).Debug("dropping outbound packet",
"fwPacket", fwPacket,
@@ -119,127 +86,8 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Parse
}
}
func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, seg, scratch, nb []byte) []byte {
if noiseutil.EncryptLockNeeded {
ci.writeLock.Lock()
}
c := ci.messageCounter.Add(1)
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
out, encErr := ci.eKey.EncryptDanger(out, out, seg, c, nb)
if noiseutil.EncryptLockNeeded {
ci.writeLock.Unlock()
}
if encErr != nil {
hostinfo.logger(f.l).Error("Failed to encrypt outgoing packet",
"error", encErr,
"udpAddr", hostinfo.GetRemote(),
"counter", c,
)
// Skip this segment; the rest of the superpacket can still go out. TCP will retransmit anything we drop here.
return nil
}
return out
}
// sendInsideMessage encrypts a firewall-approved inside packet (or every
// segment of a TSO/USO superpacket) into the caller's batch slot for
// later sendmmsg flush. Segmentation is fused with encryption here so the
// kernel-supplied superpacket bytes never get written into a separate
// scratch arena: SegmentSuperpacket builds each segment's plaintext in
// segScratch[:segLen] in turn, and we encrypt directly into a fresh SendBatch slot.
func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []byte, sendBatch *batch.SendBatch) {
ci := hostinfo.ConnectionState
if ci.eKey == nil {
return
}
// One traffic-out mark covers every segment of the superpacket; doing it
// per segment in sendInsideEncrypt paid an atomic store up to ~45 extra
// times per TSO packet, inside writeLock under boring crypto.
f.connectionManager.Out(hostinfo)
remote := hostinfo.GetRemote()
if hostinfo.lastRebindCount != f.rebindCount {
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
f.lightHouse.QueryServer(hostinfo.vpnAddrs[0])
hostinfo.lastRebindCount = f.rebindCount
if f.l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(f.l).Debug("Lighthouse update triggered for punch due to rebind counter",
"vpnAddrs", hostinfo.vpnAddrs,
)
}
}
if !remote.IsValid() { //the relay path
//first, find our relay hostinfo:
var relayHostInfo *HostInfo
var relay *Relay
var err error
for _, relayIP := range hostinfo.relayState.CopyRelayIps() {
relayHostInfo, relay, err = f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relayIP)
if err != nil {
hostinfo.relayState.DeleteRelay(relayIP)
hostinfo.logger(f.l).Info("sendNoMetrics failed to find HostInfo",
"relay", relayIP,
"error", err,
)
continue
}
break
}
if relayHostInfo == nil || relay == nil {
//failure already logged
return
}
err = tio.SegmentSuperpacket(pkt, func(seg []byte) error {
//relay header + header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305) + relay tag
scratch := sendBatch.Reserve(header.Len + header.Len + len(seg) + 16 + 16)
innerPacket := f.sendInsideEncrypt(hostinfo, ci, seg, scratch[header.Len:], nb)
if innerPacket == nil {
return nil
}
//now we need to do a relay-encrypt:
toSend, err := f.prepareSendVia(relayHostInfo, relay, innerPacket, nb, scratch, true)
if err != nil {
//already logged
return nil
}
sendBatch.Commit(toSend, relayHostInfo.GetRemote())
return nil
})
if err != nil {
hostinfo.logger(f.l).Error("Failed to segment superpacket for relay send", "error", err)
}
return
}
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
// header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305)
scratch := sendBatch.Reserve(header.Len + len(seg) + 16)
out := f.sendInsideEncrypt(hostinfo, ci, seg, scratch, nb)
if out == nil {
return nil
}
sendBatch.Commit(out, remote)
return nil
})
if err != nil {
hostinfo.logger(f.l).Error("Failed to segment superpacket for send", "error", err)
}
}
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
if !f.firewall.OutboundSendReject {
if !f.firewall.InSendReject {
return
}
@@ -248,36 +96,33 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
return
}
_, err := f.queues[q].Write(out)
_, err := f.readers[q].Write(out)
if err != nil {
f.l.Error("Failed to write to tun", "error", err)
}
}
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, rejectBuf []byte, q int) {
if !f.firewall.InboundSendReject {
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) {
if !f.firewall.OutSendReject {
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 +216,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 +224,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",
@@ -430,13 +275,21 @@ func (f *Interface) sendTo(t header.MessageType, st header.MessageSubType, ci *C
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
}
func (f *Interface) prepareSendVia(via *HostInfo,
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
// to the payload for the ultimate target host, making this a useful method for sending
// handshake messages to peers through relay tunnels.
// via is the HostInfo through which the message is relayed.
// ad is the plaintext data to authenticate, but not encrypt
// nb is a buffer used to store the nonce value, re-used for performance reasons.
// out is a buffer used to store the result of the Encrypt operation
// q indicates which writer to use to send the packet.
func (f *Interface) SendVia(via *HostInfo,
relay *Relay,
ad,
nb,
out []byte,
nocopy bool,
) ([]byte, error) {
) {
if noiseutil.EncryptLockNeeded {
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
via.ConnectionState.writeLock.Lock()
@@ -458,7 +311,7 @@ func (f *Interface) prepareSendVia(via *HostInfo,
"headerLen", len(out),
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
)
return nil, io.ErrShortBuffer
return
}
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
@@ -478,31 +331,13 @@ func (f *Interface) prepareSendVia(via *HostInfo,
}
if err != nil {
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
return nil, err
}
f.connectionManager.RelayUsed(relay.LocalIndex)
return out, nil
}
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
// to the payload for the ultimate target host, making this a useful method for sending
// handshake messages to peers through relay tunnels.
// via is the HostInfo through which the message is relayed.
// ad is the plaintext data to authenticate, but not encrypt
// nb is a buffer used to store the nonce value, re-used for performance reasons.
// out is a buffer used to store the result of the Encrypt operation
// q indicates which writer to use to send the packet.
func (f *Interface) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy)
if err != nil {
// already logged by prepareSendVia
return
}
err = f.writers[q].WriteTo(toSend, via.GetRemote())
err = f.writers[0].WriteTo(out, via.GetRemote())
if err != nil {
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
}
f.connectionManager.RelayUsed(relay.LocalIndex)
}
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
@@ -573,7 +408,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 +423,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
}
}
+58 -194
View File
@@ -4,9 +4,9 @@ import (
"context"
"errors"
"fmt"
"io"
"log/slog"
"net/netip"
"runtime"
"slices"
"sync"
"sync/atomic"
@@ -14,15 +14,12 @@ import (
"github.com/gaissmai/bart"
"github.com/rcrowley/go-metrics"
"github.com/slackhq/nebula/util"
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/overlay"
"github.com/slackhq/nebula/overlay/batch"
"github.com/slackhq/nebula/overlay/tio"
"github.com/slackhq/nebula/udp"
)
@@ -52,19 +49,7 @@ type InterfaceConfig struct {
reQueryWait time.Duration
ConntrackCacheTimeout time.Duration
// CpuAffinity, when non-empty, names the CPUs each TUN reader goroutine
// should pin to. Queue i pins to CpuAffinity[i % len(CpuAffinity)] —
// shorter lists than `routines` cycle. Empty list keeps the default
// pin-to-(i % NumCPU) behavior. Only consulted when PinThreads is true.
CpuAffinity []int
// PinThreads controls whether each TUN reader OS thread is pinned to a
// single CPU (via tun.pin_threads, default true). Pinning keeps each
// goroutine's sendmmsg on one XPS-selected NIC TX ring so per-flow
// packets stay ordered on the wire.
PinThreads bool
l *slog.Logger
l *slog.Logger
}
type Interface struct {
@@ -88,16 +73,7 @@ type Interface struct {
routines int
disconnectInvalid atomic.Bool
closed atomic.Bool
// cpuAffinity, when non-empty, names the CPUs each TUN reader goroutine
// should pin to. Queue i pins to cpuAffinity[i % len(cpuAffinity)].
// Empty falls back to the default pin-to-(allowed CPU) behavior.
// Only consulted when pinThreads is true.
cpuAffinity []int
// pinThreads controls whether listenIn pins each TUN reader OS thread to
// a CPU at all (tun.pin_threads, default true). When false, threads are
// left free to migrate as on stock nebula.
pinThreads bool
relayManager *relayManager
relayManager *relayManager
tryPromoteEvery atomic.Uint32
reQueryEvery atomic.Uint32
@@ -114,14 +90,8 @@ 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
wg sync.WaitGroup
readers []io.ReadWriteCloser
wg sync.WaitGroup
// fatalErr holds the first unexpected reader error that caused shutdown.
// nil means "no fatal error" (yet)
@@ -132,13 +102,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 +172,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 +189,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),
readers: make([]io.ReadWriteCloser, c.routines),
myVpnNetworks: cs.myVpnNetworks,
myVpnNetworksTable: cs.myVpnNetworksTable,
myVpnAddrs: cs.myVpnAddrs,
@@ -227,11 +198,8 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
relayManager: c.relayManager,
connectionManager: c.connectionManager,
conntrackCacheTimeout: c.ConntrackCacheTimeout,
cpuAffinity: c.CpuAffinity,
pinThreads: c.PinThreads,
metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)),
metricTxDropped: metrics.GetOrRegisterCounter("udp.tx.dropped", nil),
messageMetrics: c.MessageMetrics,
cachedPacketMetrics: &cachedPacketMetrics{
sent: metrics.GetOrRegisterCounter("hostinfo.cached_packets.sent", nil),
@@ -247,9 +215,6 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
ifce.connectionManager.intf = ifce
// Held until Close so waiting on the interface blocks until the resources are actually released
ifce.wg.Add(1)
return ifce, nil
}
@@ -272,48 +237,38 @@ func (f *Interface) activate() error {
"boringcrypto", boringEnabled(),
)
if f.routines > 1 && !f.outside.SupportsMultipleReaders() {
f.routines = 1
f.l.Warn("multiple udp readers are not supported on this platform, falling back to a single routine")
if f.routines > 1 {
if !f.inside.SupportsMultiqueue() || !f.outside.SupportsMultipleReaders() {
f.routines = 1
f.l.Warn("routines is not supported on this platform, falling back to a single routine")
}
}
// Prepare the tun queues. A device that can't open that many hands back
// fewer (a single queue on platforms without multiqueue support) and we
// size the reader routines to what we actually got.
queues, err := f.inside.Queues(f.routines)
if err != nil {
return err
}
if len(queues) < f.routines {
// TODO: this clamp is only safe because it is unreachable when the
// udp side has multiple readers (linux Queues opens exactly n or
// errors; every other platform already clamped routines to 1 above).
// If a platform ever returns fewer queues than routines with
// SO_REUSEPORT sockets already bound, the surplus sockets get no
// listenOut and the kernel blackholes every flow it hashes to them —
// fail loudly or close the extra sockets instead.
f.l.Warn("tun multiqueue is not supported on this platform, falling back to fewer routines",
"requested", f.routines, "opened", len(queues))
f.routines = len(queues)
}
f.queues = queues
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
for i := range f.queues {
f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l)
// Prepare n tun queues
var reader io.ReadWriteCloser = f.inside
for i := 0; i < f.routines; i++ {
if i > 0 {
reader, err = f.inside.NewMultiQueueReader()
if err != nil {
return err
}
}
f.readers[i] = reader
}
// On error the caller owns the cleanup, Control.Start cancels the service context
// before releasing our resources so a waiter never observes a live context
f.wg.Add(1) // for us to wait on Close() to return
if err = f.inside.Activate(); err != nil {
f.wg.Done()
f.inside.Close()
return err
}
return nil
}
func (f *Interface) run() {
func (f *Interface) run() (func() error, error) {
// Launch n queues to read packets from udp
for i := 0; i < f.routines; i++ {
f.wg.Go(func() {
@@ -324,18 +279,17 @@ func (f *Interface) run() {
// Launch n queues to read packets from tun dev
for i := 0; i < f.routines; i++ {
f.wg.Go(func() {
f.listenIn(f.queues[i], i)
f.listenIn(f.readers[i], i)
})
}
}
func (f *Interface) wait() error {
f.wg.Wait()
if e := f.fatalErr.Load(); e != nil {
return *e
}
return nil
return func() error {
f.wg.Wait()
if e := f.fatalErr.Load(); e != nil {
return *e
}
return nil
}, nil
}
// onFatal stores the first fatal reader error, and calls triggerShutdown if it was the first one
@@ -349,31 +303,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,25 +311,18 @@ 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()
plaintext := make([]byte, udp.MTU)
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)
}
err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
})
flusher := func() {
if err := f.batchers[i].Flush(); err != nil {
f.l.Error("Failed to flush tun coalescer", "error", err)
}
clear(rxc.hostmapCache)
}
err := li.ListenOut(listener, flusher)
// An error after teardown began is shutdown noise, the closed flag covers resources
// Close releases itself and the cancelled ctx covers ones torn down by their owners
// reacting to it, like the user device pipes
if err != nil && !f.closed.Load() && f.ctx.Err() == nil {
if err != nil && !f.closed.Load() {
f.l.Error("Error while reading inbound packet, closing", "error", err)
f.onFatal(err)
}
@@ -408,80 +330,30 @@ 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)
}
rejectBuf := make([]byte, mtu)
arenaSize := batch.SendBatchCap * (udp.MTU + 32)
sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, arenaSize)
fwPacket := &firewall.ParsedPacket{}
func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
packet := make([]byte, mtu)
out := make([]byte, mtu)
fwPacket := &firewall.Packet{}
nb := make([]byte, 12, 12)
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
for {
pkts, err := queue.Read()
n, err := reader.Read(packet)
if err != nil {
// Same shutdown noise handling as listenOut
if !f.closed.Load() && f.ctx.Err() == nil {
if !f.closed.Load() {
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i)
f.onFatal(err)
}
break
}
for _, pkt := range pkts {
f.consumeInsidePacket(pkt, fwPacket, nb, sb, rejectBuf, i, conntrackCache.Get())
// Flush incrementally once a full sendmmsg batch has
// accumulated so the first packets of a deep read drain
// hit the wire while the rest are still being encrypted.
if sb.Len() >= batch.SendBatchCap {
f.flushSendBatch(sb, i)
}
}
f.flushSendBatch(sb, i)
f.consumeInsidePacket(packet[:n], fwPacket, nb, out, i, conntrackCache.Get())
}
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)
@@ -670,15 +542,9 @@ func (f *Interface) GetCertState() *CertState {
return f.pki.getCertState()
}
// Close releases the interface's resources: the udp sockets and the tun device.
// It is idempotent and safe to call at any point in the lifecycle, including on an interface that never activated,
// calls after the first return nil without doing anything.
func (f *Interface) Close() error {
if !f.closed.CompareAndSwap(false, true) {
return nil
}
var errs []error
f.closed.Store(true)
// Release the udp readers
for i, u := range f.writers {
@@ -694,8 +560,6 @@ func (f *Interface) Close() error {
if closeErr != nil {
errs = append(errs, closeErr)
}
// Release the construction token so waiters know the resources are gone
f.wg.Done()
return errors.Join(errs...)
}
+3 -27
View File
@@ -199,7 +199,7 @@ func ipv4CreateRejectTCPPacket(packet []byte, out []byte) []byte {
}
func ipv6CreateRejectPacket(packet []byte, out []byte) []byte {
proto, offset, isFragment := IPv6FindUpperProtocol(packet)
proto, offset, isFragment := ipv6FindUpperProtocol(packet)
if isFragment {
return nil
}
@@ -333,34 +333,11 @@ func ipv6CreateRejectTCPPacket(packet []byte, out []byte, offset int) []byte {
return out
}
// maxIPv6ExtHeaders caps the extension-header walk in IPv6FindUpperProtocol.
// RFC 8200 legal chains are shorter (each header at most once, Destination
// Options at most twice), so the cap only bites crafted packets, which would
// otherwise make us walk their whole payload 8 bytes at a time.
const maxIPv6ExtHeaders = 8
// IPv6FindUpperProtocol walks packet's IPv6 extension-header chain and
// returns the terminal (upper-layer) protocol number, the byte offset where
// that protocol's header begins, and whether the packet is a non-first
// fragment. It steps over Hop-by-Hop (0), Routing (43), Fragment (44),
// AH (51), and Destination Options (60); anything else — including ESP,
// whose payload is encrypted — terminates the walk.
//
// For a non-first fragment, nextHeader still names the flow's upper
// protocol (copied from the fragment header) but offset points at fragment
// payload, not a real transport header: consult isFragment before
// dereferencing offset. If the chain is truncated, over-long, or the packet
// is shorter than an IPv6 header, the walk stops early and nextHeader is
// whatever it stopped on (59, IPPROTO_NONE, for the too-short case) —
// callers treat any non-transport result as unclassifiable.
func IPv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool) {
if len(packet) < ipv6.HeaderLen {
return 59, 0, false // IPPROTO_NONE: nothing to classify
}
func ipv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool) {
nextHeader = packet[6]
offset = ipv6.HeaderLen
for range maxIPv6ExtHeaders {
for {
switch nextHeader {
case 0, 43, 60: // Hop-by-Hop, Routing, Destination
if len(packet) < offset+2 {
@@ -390,7 +367,6 @@ func IPv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragm
return nextHeader, offset, isFragment
}
}
return nextHeader, offset, isFragment
}
func CreateICMPEchoResponse(packet, out []byte) []byte {
-159
View File
@@ -1,7 +1,6 @@
package iputil
import (
"bytes"
"encoding/binary"
"net"
"testing"
@@ -180,46 +179,6 @@ func Test_CreateRejectPacket_NoICMPError(t *testing.T) {
}
}
// Test_CreateRejectPacket_RespectsCap ensures it is impossible for
// an oversized ICMPv6 reject to overwrite the neighbor segment's bytes.
func Test_CreateRejectPacket_RespectsCap(t *testing.T) {
src := net.ParseIP("fd00::1")
dst := net.ParseIP("fd00::2")
// Inner IPv6 UDP packet. An ICMPv6 reject copies the whole inner packet
// plus a 48-byte header (40 IPv6 + 8 ICMPv6), so it needs 48 more bytes
// than the inner packet length.
inner := makeIPv6Packet(src, dst, 17, make([]byte, 20))
// The ciphertext scratch reused as the reject buffer is the received
// datagram: 16-byte Nebula header + inner + 16-byte AEAD tag. That is only
// 32 bytes of slack, so a full ICMPv6 reject overruns it by 16 bytes.
const nebulaOverhead = 32
segLen := len(inner) + nebulaOverhead
// Shared backing row laid out as [segment][neighbor's 16-byte Nebula header].
const neighborHdr = 16
sentinel := bytes.Repeat([]byte{0xAB}, neighborHdr)
// Uncapped: the slice's capacity reaches into the neighbor, reproducing
// the overrun that silently drops the neighbor packet.
backing := make([]byte, segLen+neighborHdr)
copy(backing[segLen:], sentinel)
reject := CreateRejectPacket(inner, backing[:segLen])
assert.NotNil(t, reject, "uncapped buffer reaches into the neighbor, so the reject is built")
assert.NotEqual(t, sentinel, backing[segLen:segLen+neighborHdr],
"without the cap the oversized reject overruns into the neighbor segment")
// Capped (the fix): cap==len, so the builder cannot exceed the segment. The
// reject does not fit, so it is refused rather than corrupting the neighbor.
backing = make([]byte, segLen+neighborHdr)
copy(backing[segLen:], sentinel)
reject = CreateRejectPacket(inner, backing[:segLen:segLen])
assert.Nil(t, reject, "capped segment is 16 bytes too small for a full ICMPv6 reject, so it is refused")
assert.Equal(t, sentinel, backing[segLen:segLen+neighborHdr],
"capped segment must leave the neighbor untouched")
}
func makeIPv6Packet(src, dst net.IP, nextHeader uint8, payload []byte) []byte {
b := make([]byte, ipv6.HeaderLen+len(payload))
b[0] = ipv6.Version << 4
@@ -515,121 +474,3 @@ func TestCreateICMPEchoResponse_IPv6_NotICMPv6(t *testing.T) {
result := CreateICMPEchoResponse(packet, out)
assert.Nil(t, result)
}
func TestIPv6FindUpperProtocol(t *testing.T) {
src := net.ParseIP("fd00::1")
dst := net.ParseIP("fd00::2")
// extHdr builds one 8-byte-unit extension header: next, hdrExtLen
// ((extra+1)*8 bytes total), padded to size.
extHdr := func(next uint8, extra int) []byte {
b := make([]byte, (extra+1)*8)
b[0] = next
b[1] = uint8(extra)
return b
}
t.Run("no extension headers", func(t *testing.T) {
for _, proto := range []uint8{6, 17, 58} {
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, proto, make([]byte, 20)))
assert.Equal(t, proto, nh)
assert.Equal(t, ipv6.HeaderLen, offset)
assert.False(t, frag)
}
})
t.Run("hop-by-hop then TCP", func(t *testing.T) {
payload := append(extHdr(6, 0), make([]byte, 20)...)
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, payload))
assert.Equal(t, uint8(6), nh)
assert.Equal(t, ipv6.HeaderLen+8, offset)
assert.False(t, frag)
})
t.Run("chained headers honor length units", func(t *testing.T) {
// Hop-by-Hop (8B) -> Dest Options (16B) -> Routing (8B) -> UDP.
payload := extHdr(60, 0)
payload = append(payload, extHdr(43, 1)...)
payload = append(payload, extHdr(17, 0)...)
payload = append(payload, make([]byte, 8)...)
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, payload))
assert.Equal(t, uint8(17), nh)
assert.Equal(t, ipv6.HeaderLen+8+16+8, offset)
assert.False(t, frag)
})
t.Run("AH length is in 4-byte units plus 2", func(t *testing.T) {
// AH payload-len byte 4 -> (4+2)*4 = 24 bytes on the wire.
ah := make([]byte, 24)
ah[0] = 6
ah[1] = 4
payload := append(ah, make([]byte, 20)...)
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 51, payload))
assert.Equal(t, uint8(6), nh)
assert.Equal(t, ipv6.HeaderLen+24, offset)
assert.False(t, frag)
})
t.Run("first fragment walks to the transport header", func(t *testing.T) {
frag := make([]byte, 8)
frag[0] = 17
binary.BigEndian.PutUint16(frag[2:4], 0x0001) // offset 0, M=1
payload := append(frag, make([]byte, 8)...)
nh, offset, isFrag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 44, payload))
assert.Equal(t, uint8(17), nh)
assert.Equal(t, ipv6.HeaderLen+8, offset)
assert.False(t, isFrag, "first fragment carries the real transport header")
})
t.Run("non-first fragment is flagged", func(t *testing.T) {
frag := make([]byte, 8)
frag[0] = 17
binary.BigEndian.PutUint16(frag[2:4], 1<<3) // offset 1, M=0
payload := append(frag, make([]byte, 8)...)
nh, _, isFrag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 44, payload))
assert.Equal(t, uint8(17), nh, "fragment header still names the flow's L4")
assert.True(t, isFrag, "offset points at fragment payload, not a header")
})
t.Run("ESP terminates the walk", func(t *testing.T) {
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 50, make([]byte, 16)))
assert.Equal(t, uint8(50), nh)
assert.Equal(t, ipv6.HeaderLen, offset)
assert.False(t, frag)
})
t.Run("unknown protocol terminates the walk", func(t *testing.T) {
nh, offset, _ := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 132, make([]byte, 16))) // SCTP
assert.Equal(t, uint8(132), nh)
assert.Equal(t, ipv6.HeaderLen, offset)
})
t.Run("truncated extension header stops the walk", func(t *testing.T) {
// Next header says Hop-by-Hop but the packet ends at the IPv6 header.
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, nil))
assert.Equal(t, uint8(0), nh, "unresolvable chain returns the extension header it stopped on")
assert.Equal(t, ipv6.HeaderLen, offset)
assert.False(t, frag)
})
t.Run("crafted over-long chain hits the cap", func(t *testing.T) {
// Ten chained Hop-by-Hop headers, then TCP. Illegal per RFC 8200
// (Hop-by-Hop may only appear first); the cap must stop the walk
// before it resolves rather than crawling arbitrary crafted chains.
var payload []byte
for i := 0; i < 9; i++ {
payload = append(payload, extHdr(0, 0)...)
}
payload = append(payload, extHdr(6, 0)...)
payload = append(payload, make([]byte, 20)...)
nh, _, _ := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, payload))
assert.Equal(t, uint8(0), nh, "walk must stop at the cap, not resolve to TCP")
})
t.Run("packet shorter than an IPv6 header", func(t *testing.T) {
nh, offset, frag := IPv6FindUpperProtocol(make([]byte, 39))
assert.Equal(t, uint8(59), nh) // IPPROTO_NONE
assert.Equal(t, 0, offset)
assert.False(t, frag)
})
}
+1 -9
View File
@@ -36,10 +36,6 @@ type LightHouse struct {
myVpnNetworksTable *bart.Lite
punchy *Punchy
// localAddrsFn enumerates the underlay addresses we advertise. It is a field so tests can supply simulated
// addresses rather than whatever this machine's NICs happen to be. Set it before Start.
localAddrsFn func(*LocalAllowList) []netip.Addr
// Local cache of answers from light houses
// map of vpn addr to answers
addrMap map[netip.Addr]*RemoteList
@@ -111,10 +107,6 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
l: l,
}
h.localAddrsFn = func(al *LocalAllowList) []netip.Addr {
return localAddrs(h.l, al)
}
lighthouses := make([]netip.Addr, 0)
h.lighthouses.Store(&lighthouses)
staticList := make(map[netip.Addr]struct{})
@@ -926,7 +918,7 @@ func (lh *LightHouse) SendUpdate() {
}
lal := lh.GetLocalAllowList()
for _, e := range lh.localAddrsFn(lal) {
for _, e := range localAddrs(lh.l, lal) {
if lh.myVpnNetworksTable.Contains(e) {
continue
}
+1 -1
View File
@@ -498,7 +498,7 @@ type testEncWriter struct {
protocolVersion cert.Version
}
func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool) {
}
func (tw *testEncWriter) Handshake(vpnIp netip.Addr) {
}
+1 -108
View File
@@ -6,14 +6,11 @@ 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"
@@ -36,9 +33,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
buildVersion = moduleVersion()
}
// Debug builds (-tags debug) serve pprof on :6060; a no-op otherwise.
startPprofServer(ctx, l)
// Print the config if in test, the exit comes later
if configTest {
b, err := yaml.Marshal(c.Settings)
@@ -136,17 +130,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
udpConns := make([]udp.Conn, routines)
port := c.GetInt("listen.port", 0)
// Callers get no handle to these until the Control is returned, release them on any error.
defer func() {
if reterr != nil {
for _, u := range udpConns {
if u != nil {
_ = u.Close()
}
}
}
}()
if !configTest {
rawListenHost := c.GetString("listen.host", "0.0.0.0")
var listenHost netip.Addr
@@ -167,13 +150,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)
}
@@ -222,21 +199,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
l.Warn("Failed to start DNS responder", "error", err)
}
pinThreads := c.GetBool("tun.pin_threads", true)
cpuAffinity := parseCpuAffinity(c, l, routines)
if pinThreads && routines > 1 && len(cpuAffinity) == 0 && !configTest {
// The operator didn't choose pin CPUs, so pick a default set that
// prefers performance cores and doesn't stack co-located instances
// onto allowed[0]. The bound UDP port keys the per-instance spread:
// distinct across instances sharing a box, stable across restarts.
// A nil result keeps listenIn's stock allowed[i] fallback.
key := uint64(os.Getpid())
if ap, err := udpConns[0].LocalAddr(); err == nil && ap.Port() != 0 {
key = uint64(ap.Port())
}
cpuAffinity = cpupick.Default(routines, key, l)
}
ifConfig := &InterfaceConfig{
HostMap: hostMap,
Inside: tun,
@@ -258,8 +220,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
relayManager: NewRelayManager(ctx, l, hostMap, c),
punchy: punchy,
ConntrackCacheTimeout: conntrackCacheTimeout,
CpuAffinity: cpuAffinity,
PinThreads: pinThreads,
l: l,
}
@@ -297,8 +257,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,75 +267,10 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
statsStart: stats.Start,
dnsStart: ds.Start,
lighthouseStart: lightHouse.StartUpdateWorker,
networkChangeStart: networkChanges.Start,
connectionManagerStart: connManager.Start,
}, nil
}
// parseCpuAffinity reads `tun.cpu_affinity` from the config — a list of
// integer CPU IDs, one per TUN reader goroutine. Empty / unset returns nil
// (listenIn falls back to spreading queues across the allowed CPU set).
// Length mismatch with `routines` is a warning, not an error: shorter lists
// are modulo-cycled across queues, longer lists' tail is ignored. Invalid
// entries (non-integer, or a CPU ID we're not allowed to run on) are also a
// warning and disable the override entirely so we don't silently pin to the
// wrong CPU. Entries are validated against the process's current affinity
// mask (util.AllowedCPUs) rather than 0..NumCPU-1: under a cgroup cpuset or
// taskset the runnable IDs are frequently not that contiguous range, and
// pinning to an unrunnable ID always fails. If the allowed set can't be
// determined we fall back to a plain non-negative check.
func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int {
raw := c.Get("tun.cpu_affinity")
if raw == nil {
return nil
}
rv, ok := raw.([]any)
if !ok {
l.Warn("tun.cpu_affinity must be a list of integers; ignoring", "value", raw)
return nil
}
// allowed is the set of CPU IDs we're actually permitted to run on. A nil
// slice (unsupported platform or lookup error) means "can't tell", so we
// only apply the weaker non-negative check in that case.
allowed, err := util.AllowedCPUs()
if err != nil {
l.Warn("could not determine allowed CPUs; validating tun.cpu_affinity against non-negative only", "error", err)
allowed = nil
}
cpus := make([]int, 0, len(rv))
for i, e := range rv {
var cpu int
switch v := e.(type) {
case int:
cpu = v
case int64:
cpu = int(v)
case float64:
cpu = int(v)
default:
l.Warn("tun.cpu_affinity entry not an integer; ignoring affinity",
"index", i, "value", e)
return nil
}
if cpu < 0 {
l.Warn("tun.cpu_affinity entry out of range; ignoring affinity",
"index", i, "cpu", cpu)
return nil
}
if len(allowed) > 0 && !slices.Contains(allowed, cpu) {
l.Warn("tun.cpu_affinity entry not in allowed CPU set; ignoring affinity",
"index", i, "cpu", cpu, "allowed", allowed)
return nil
}
cpus = append(cpus, cpu)
}
if len(cpus) != routines {
l.Warn("tun.cpu_affinity length doesn't match routines; queues will modulo-cycle through the list",
"affinity_len", len(cpus), "routines", routines)
}
return cpus
}
func moduleVersion() string {
info, ok := debug.ReadBuildInfo()
if !ok {
-51
View File
@@ -1,51 +0,0 @@
package nebula
import (
"testing"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/test"
"github.com/slackhq/nebula/util"
"github.com/stretchr/testify/assert"
)
func TestParseCpuAffinity(t *testing.T) {
l := test.NewLogger()
// newConfig returns a config.C with tun.cpu_affinity set to v. A nil v
// leaves the key unset.
newConfig := func(v any) *config.C {
c := config.NewC(l)
if v != nil {
c.Settings["tun"] = map[string]any{"cpu_affinity": v}
}
return c
}
// unset -> nil (listenIn falls back to spreading across the allowed set)
assert.Nil(t, parseCpuAffinity(newConfig(nil), l, 1))
// Pick a CPU we're actually allowed to run on so a valid list survives
// validation regardless of the host's affinity mask.
allowed, _ := util.AllowedCPUs()
validCPU := 0
if len(allowed) > 0 {
validCPU = allowed[0]
}
// valid list -> parsed through unchanged
assert.Equal(t, []int{validCPU, validCPU}, parseCpuAffinity(newConfig([]any{validCPU, validCPU}), l, 2))
// a negative entry is out of range on every platform -> disables the override
assert.Nil(t, parseCpuAffinity(newConfig([]any{validCPU, -1}), l, 2))
// a non-integer entry -> disables the override
assert.Nil(t, parseCpuAffinity(newConfig([]any{validCPU, "not-a-cpu"}), l, 2))
// a CPU id outside the allowed set -> disables the override. Only assertable
// where we can enumerate the allowed set (e.g. linux); 1<<20 is far beyond
// any representable CPU id so it can never be in the mask.
if len(allowed) > 0 {
assert.Nil(t, parseCpuAffinity(newConfig([]any{1 << 20}), l, 1))
}
}
-45
View File
@@ -164,48 +164,3 @@ func TestCipherStateNilSafety(t *testing.T) {
assert.Empty(t, out)
assert.Equal(t, 0, cc.Overhead())
}
func TestCipherStateAESGCMInPlaceDecrypt(t *testing.T) {
enc, dec := buildCipherStates(t, CipherAESGCM)
inPlaceDecrypt(t, NewCipherStateAESGCM(enc), NewCipherStateAESGCM(dec))
}
func TestCipherStateChaChaPolyInPlaceDecrypt(t *testing.T) {
enc, dec := buildCipherStates(t, noise.CipherChaChaPoly)
inPlaceDecrypt(t, NewCipherStateChaChaPoly(enc), NewCipherStateChaChaPoly(dec))
}
func inPlaceDecrypt(t *testing.T, enc, dec CipherState) {
t.Helper()
const hdrLen = 16
plaintext := []byte("in-place decrypt should replace the ciphertext bytes")
nb := make([]byte, 12)
// packet = [16-byte header | ciphertext+tag], like a nebula Message.
packet := make([]byte, hdrLen, hdrLen+len(plaintext)+enc.Overhead())
for i := range packet {
packet[i] = byte(i)
}
packet, err := enc.EncryptDanger(packet, packet[:hdrLen], plaintext, 1, nb)
require.NoError(t, err)
// Simulate a GRO row: [packet | next segment]. A failed auth on packet
// may zero packet's plaintext region but must not touch the header, the
// tag, or the neighboring segment.
neighbor := []byte("next coalesced segment, must stay intact")
row := append(append([]byte(nil), packet...), neighbor...)
tampered := row[:len(packet)]
tampered[hdrLen] ^= 0x01
_, err = dec.DecryptDanger(tampered[hdrLen:hdrLen], tampered[:hdrLen], tampered[hdrLen:], 1, nb)
require.Error(t, err)
assert.Equal(t, packet[:hdrLen], tampered[:hdrLen], "failed auth must not touch the header")
assert.Equal(t, packet[len(packet)-dec.Overhead():], tampered[len(tampered)-dec.Overhead():],
"failed auth must not touch the tag")
assert.Equal(t, neighbor, row[len(packet):], "failed auth must not touch the next segment")
out, err := dec.DecryptDanger(packet[hdrLen:hdrLen], packet[:hdrLen], packet[hdrLen:], 1, nb)
require.NoError(t, err)
assert.Equal(t, plaintext, out)
// The plaintext must be IN the packet buffer, not a fresh allocation.
assert.Equal(t, &packet[hdrLen], &out[0], "plaintext must alias the packet buffer")
}
+83 -76
View File
@@ -13,7 +13,6 @@ import (
"github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/overlay/batch"
"golang.org/x/net/ipv4"
)
@@ -23,11 +22,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) {
err := h.Parse(packet)
if err != nil {
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
@@ -95,7 +90,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 +102,27 @@ 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)
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)
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) {
// 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)
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
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,45 @@ 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
}
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) {
err := newPacket(out, true, fwPacket)
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
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, 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.readers[q].Write(out)
if err != nil {
f.l.Error("Failed to write to tun", "error", err)
}
+5 -90
View File
@@ -17,7 +17,7 @@ import (
)
func Test_newPacket(t *testing.T) {
p := &firewall.ParsedPacket{}
p := &firewall.Packet{}
// length fails
err := newPacket([]byte{}, true, p)
@@ -96,7 +96,7 @@ func Test_newPacket(t *testing.T) {
}
func Test_newPacket_v6(t *testing.T) {
p := &firewall.ParsedPacket{}
p := &firewall.Packet{}
// invalid ipv6
ip := layers.IPv6{
@@ -345,7 +345,7 @@ func Test_newPacket_v6(t *testing.T) {
}
func Test_newPacket_ipv6Fragment(t *testing.T) {
p := &firewall.ParsedPacket{}
p := &firewall.Packet{}
ip := &layers.IPv6{
Version: 6,
@@ -525,7 +525,7 @@ func BenchmarkParseV6(b *testing.B) {
secondFrag = append(secondFrag, fragHeader...)
secondFrag = append(secondFrag, []byte{0xde, 0xad, 0xbe, 0xef}...)
fp := &firewall.ParsedPacket{}
fp := &firewall.Packet{}
b.Run("Normal", func(b *testing.B) {
for i := 0; i < b.N; i++ {
@@ -649,7 +649,7 @@ func serializeAH(ah *layers.IPSecAH) []byte {
// host OS parses the real header, a firewall port/proto bypass. The fix makes parseV6 land
// on the same offset the host does.
func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
p := &firewall.ParsedPacket{}
p := &firewall.Packet{}
const (
hdrLen = 40 // IPv6 header
@@ -675,88 +675,3 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
// the host delivers to, not the forged 443 at the overflowed offset.
assert.Equal(t, uint16(22), p.LocalPort, "firewall must parse the real transport header, not the overflowed offset")
}
// Test_newPacket_parsedFields pins the ParsedPacket byproducts the RX
// batcher consumes: IPHdrLen (the true L4 offset) and FragAny (any fragment
// shape at all — unlike Packet.Fragment, which is port-oriented and true
// only for non-first fragments).
func Test_newPacket_parsedFields(t *testing.T) {
p := &firewall.ParsedPacket{}
// Plain IPv4 TCP, IHL 20: L4 offset 20, no fragment shape.
v4 := make([]byte, 28)
v4[0] = 0x45
v4[9] = firewall.ProtoTCP
binary.BigEndian.PutUint16(v4[6:8], 0x4000) // DF only
require.NoError(t, newPacket(v4, true, p))
assert.Equal(t, 20, p.IPHdrLen)
assert.False(t, p.FragAny)
assert.False(t, p.Fragment)
// IPv4 first fragment (MF set, offset 0): the firewall can read ports
// (Fragment false) but the coalescer must not touch it (FragAny true).
ff := make([]byte, 28)
ff[0] = 0x45
ff[9] = firewall.ProtoUDP
binary.BigEndian.PutUint16(ff[6:8], 0x2000) // MF, offset 0
require.NoError(t, newPacket(ff, true, p))
assert.False(t, p.Fragment)
assert.True(t, p.FragAny)
assert.Equal(t, 20, p.IPHdrLen)
// IPv4 non-first fragment (nonzero offset): both flags set.
nf := make([]byte, 28)
nf[0] = 0x45
nf[9] = firewall.ProtoUDP
binary.BigEndian.PutUint16(nf[6:8], 0x00b9)
require.NoError(t, newPacket(nf, true, p))
assert.True(t, p.Fragment)
assert.True(t, p.FragAny)
// IPv4 with options (IHL 24): IPHdrLen tracks the real L4 offset.
opts := make([]byte, 32)
opts[0] = 0x46
opts[9] = firewall.ProtoTCP
binary.BigEndian.PutUint16(opts[6:8], 0x4000)
require.NoError(t, newPacket(opts, true, p))
assert.Equal(t, 24, p.IPHdrLen)
assert.False(t, p.FragAny)
// Plain IPv6 TCP: L4 at 40.
v6 := make([]byte, 60)
v6[0] = 0x60
v6[6] = firewall.ProtoTCP
require.NoError(t, newPacket(v6, true, p))
assert.Equal(t, 40, p.IPHdrLen)
assert.False(t, p.FragAny)
// IPv6 hop-by-hop then TCP: IPHdrLen lands past the extension header.
hbh := make([]byte, 60)
hbh[0] = 0x60
hbh[6] = 0 // hop-by-hop
hbh[40] = firewall.ProtoTCP
hbh[41] = 0 // HdrExtLen 0 -> 8-byte header
require.NoError(t, newPacket(hbh, true, p))
assert.Equal(t, 48, p.IPHdrLen)
assert.False(t, p.FragAny)
// IPv6 first fragment: terminal proto resolved, FragAny set, Fragment not.
f6 := make([]byte, 60)
f6[0] = 0x60
f6[6] = 44 // fragment extension header
f6[40] = firewall.ProtoUDP
require.NoError(t, newPacket(f6, true, p))
assert.True(t, p.FragAny)
assert.False(t, p.Fragment)
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
// IPv6 non-first fragment: both set, walk stops at the fragment header.
f6n := make([]byte, 60)
f6n[0] = 0x60
f6n[6] = 44
f6n[40] = firewall.ProtoUDP
binary.BigEndian.PutUint16(f6n[42:44], 0x0008)
require.NoError(t, newPacket(f6n, true, p))
assert.True(t, p.Fragment)
assert.True(t, p.FragAny)
}
-11
View File
@@ -1,11 +0,0 @@
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
}
-187
View File
@@ -1,187 +0,0 @@
package batch
import (
"encoding/binary"
"math/rand"
"testing"
)
// The checksum-seeding helpers feed the virtio NEEDS_CSUM contract: the L4
// checksum field is pre-loaded with the folded (not inverted) pseudo-header
// sum, and the kernel later adds the L4 byte sum and inverts. A wrong seed
// produces packets every receiver silently drops, with nothing failing on
// our side — so these tests check the helpers against an independent
// RFC 1071 reference built from explicit pseudo-header bytes, never against
// the production checksum code.
// refSum accumulates big-endian 16-bit words of b (odd tail zero-padded)
// into a wide one's-complement accumulator.
func refSum(b []byte) uint64 {
var s uint64
for i := 0; i+1 < len(b); i += 2 {
s += uint64(b[i])<<8 | uint64(b[i+1])
}
if len(b)%2 == 1 {
s += uint64(b[len(b)-1]) << 8
}
return s
}
// refFold folds a wide one's-complement accumulator to 16 bits.
func refFold(s uint64) uint16 {
for s>>16 != 0 {
s = s&0xffff + s>>16
}
return uint16(s)
}
func TestFoldOnceNoInvertEdgeCases(t *testing.T) {
cases := []uint32{
0, 1, 0xffff,
0x10000, // single carry
0x1fffe, // 0xffff + 0xffff: carry produces another 0xffff
0xffff0000, // high half only
0xfffeffff, // fold yields 0x1fffd: needs a second fold
0xffffffff, // worst case
0x00010001, // simple two-word
}
for _, c := range cases {
want := refFold(uint64(c))
if got := foldOnceNoInvert(c); got != want {
t.Errorf("foldOnceNoInvert(%#x) = %#x, want %#x", c, got, want)
}
// Folding a folded value must be a no-op.
if got := foldOnceNoInvert(uint32(foldOnceNoInvert(c))); got != foldOnceNoInvert(c) {
t.Errorf("foldOnceNoInvert not idempotent at %#x", c)
}
}
}
func TestPseudoSumIPv4MatchesReference(t *testing.T) {
cases := []struct {
name string
src, dst [4]byte
proto byte
l4Len int
}{
{"simple", [4]byte{10, 0, 0, 1}, [4]byte{10, 0, 0, 2}, 6, 20},
{"zero-len", [4]byte{192, 168, 1, 1}, [4]byte{192, 168, 1, 2}, 17, 0},
{"max-len", [4]byte{1, 2, 3, 4}, [4]byte{5, 6, 7, 8}, 6, 65535},
{"carry-heavy", [4]byte{255, 255, 255, 255}, [4]byte{255, 255, 255, 254}, 17, 65535},
{"broadcastish", [4]byte{255, 255, 255, 255}, [4]byte{255, 255, 255, 255}, 255, 65535},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
// RFC 793 pseudo-header: src(4) dst(4) zero(1) proto(1) len(2).
ph := make([]byte, 12)
copy(ph[0:4], c.src[:])
copy(ph[4:8], c.dst[:])
ph[9] = c.proto
binary.BigEndian.PutUint16(ph[10:12], uint16(c.l4Len))
want := refFold(refSum(ph))
got := foldOnceNoInvert(pseudoSumIPv4(c.src[:], c.dst[:], c.proto, c.l4Len))
if got != want {
t.Errorf("fold(pseudoSumIPv4) = %#x, want %#x", got, want)
}
})
}
}
func TestPseudoSumIPv6MatchesReference(t *testing.T) {
ones := func(b byte) (a [16]byte) {
for i := range a {
a[i] = b
}
return
}
cases := []struct {
name string
src, dst [16]byte
proto byte
l4Len int
}{
{"simple", [16]byte{0xfe, 0x80, 15: 1}, [16]byte{0xfe, 0x80, 15: 2}, 6, 20},
{"zero-len", [16]byte{0x20, 0x01, 15: 9}, [16]byte{0x20, 0x01, 15: 8}, 17, 0},
{"max-u16-len", ones(0xff), ones(0xfe), 6, 65535},
{"len-past-u16", ones(0xff), ones(0xff), 17, 0x12345}, // exercises the 32-bit split
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
// RFC 8200 pseudo-header: src(16) dst(16) len(4) zero(3) next(1).
ph := make([]byte, 40)
copy(ph[0:16], c.src[:])
copy(ph[16:32], c.dst[:])
binary.BigEndian.PutUint32(ph[32:36], uint32(c.l4Len))
ph[39] = c.proto
want := refFold(refSum(ph))
got := foldOnceNoInvert(pseudoSumIPv6(c.src[:], c.dst[:], c.proto, c.l4Len))
if got != want {
t.Errorf("fold(pseudoSumIPv6) = %#x, want %#x", got, want)
}
})
}
}
func TestIPv4HdrChecksumMatchesReference(t *testing.T) {
rng := rand.New(rand.NewSource(0x1791))
for _, hdrLen := range []int{20, 24, 40, 60} {
for trial := 0; trial < 200; trial++ {
hdr := make([]byte, hdrLen)
rng.Read(hdr)
hdr[0] = 0x40 | byte(hdrLen/4)
hdr[10], hdr[11] = 0, 0 // checksum field zeroed, as the contract requires
want := ^refFold(refSum(hdr))
got := ipv4HdrChecksum(hdr)
if got != want {
t.Fatalf("ipv4HdrChecksum(len=%d trial=%d) = %#x, want %#x", hdrLen, trial, got, want)
}
// Receiver-side property: with the checksum stored, the full
// header must sum to all-ones.
binary.BigEndian.PutUint16(hdr[10:12], got)
if v := refFold(refSum(hdr)); v != 0xffff {
t.Fatalf("stored checksum does not validate: full-header fold = %#x", v)
}
}
}
}
// TestChecksumSeedReceiverAcceptance is the end-to-end property the helpers
// exist for: seed the TCP checksum field with fold(pseudoSum), do what the
// kernel's NEEDS_CSUM completion does (one's-complement sum over the L4
// bytes including the seed, then invert, then store), and verify the result
// the way a receiver does (pseudo-header + L4 must sum to all-ones).
func TestChecksumSeedReceiverAcceptance(t *testing.T) {
rng := rand.New(rand.NewSource(0x1826))
for trial := 0; trial < 200; trial++ {
src := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
dst := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
payLen := rng.Intn(1500)
l4 := make([]byte, 20+payLen)
rng.Read(l4)
// Seed exactly as flushSlot does.
seed := foldOnceNoInvert(pseudoSumIPv4(src[:], dst[:], 6, len(l4)))
binary.BigEndian.PutUint16(l4[16:18], seed)
// Kernel NEEDS_CSUM completion: sum the L4 region (seed included,
// which is equivalent to summing with the field zeroed and folding
// the seed in), invert, store.
final := ^refFold(refSum(l4[:16]) + uint64(seed) + refSum(l4[18:]))
binary.BigEndian.PutUint16(l4[16:18], final)
// Receiver validation.
ph := make([]byte, 12)
copy(ph[0:4], src[:])
copy(ph[4:8], dst[:])
ph[9] = 6
binary.BigEndian.PutUint16(ph[10:12], uint16(len(l4)))
if v := refFold(refSum(ph) + refSum(l4)); v != 0xffff {
t.Fatalf("trial %d: receiver rejects packet: fold = %#x (seed=%#x final=%#x payLen=%d)",
trial, v, seed, final, payLen)
}
}
}
-161
View File
@@ -1,161 +0,0 @@
package batch
import (
"bytes"
"encoding/binary"
)
// flowKey identifies a transport flow by {src, dst, sport, dport, family}.
// Comparable, so map lookups and linear scans over the slot list stay tight.
// Shared by the TCP and UDP coalescers; each coalescer keeps its own
// openSlots map, so a TCP and UDP flow on the same 5-tuple-without-proto never alias.
type flowKey struct {
src, dst [16]byte
sport, dport uint16
isV6 bool
}
// initialSlots is the starting capacity of the slot pool.
// One flow per packet is the worst case,
// so this matches a typical carrier-side recvmmsg batch on the UDP socket.
const initialSlots = 64
// parseIPAt validates the IP header for lane parsing. newPacket already resolved the L4 protocol
// and offset for the firewall, so there is no proto sniff here; the caller's ipHdrLen is
// cross-checked instead. A plain header (v4 IHL 20, v6 exactly 40) is the only coalesceable
// shape. The v6 check is load-bearing: it rejects extension-header packets whose L4 is not at
// byte 40.
//
// The prologues fill fk's addresses and family in place (ports belong to the L4 parser; fk must
// be zero on entry so the v4 path leaves src[4:]/dst[4:] clear for map equality) and return pkt
// trimmed to the IP-declared length. The receiver-as-out-pointer shape is deliberate: these
// functions are too big to inline, and returning structs by value put five 64-byte copies on the
// per-packet path.
func (fk *flowKey) parseIPAt(pkt []byte, ipHdrLen int) ([]byte, bool) {
if len(pkt) < 20 {
return nil, false
}
switch pkt[0] >> 4 {
case 4:
if ipHdrLen != 20 {
return nil, false
}
return fk.parseIPv4Prologue(pkt)
case 6:
if ipHdrLen != 40 || len(pkt) < 40 {
return nil, false
}
return fk.parseIPv6Prologue(pkt)
}
return nil, false
}
// parseIPv4Prologue is the shared IPv4 tail of the prologue entries; the caller has verified
// len(pkt) >= 20 and the version.
func (fk *flowKey) parseIPv4Prologue(pkt []byte) ([]byte, bool) {
ihl := int(pkt[0]&0x0f) * 4
if ihl != 20 {
return nil, false
}
// Reject any fragmentation (MF or nonzero offset). The dispatcher already gated FragAny; kept
// as defense in depth, since a fragment folded into a superpacket would corrupt reassembly.
if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 {
return nil, false
}
totalLen := int(binary.BigEndian.Uint16(pkt[2:4]))
if totalLen > len(pkt) || totalLen < ihl {
return nil, false
}
fk.isV6 = false
copy(fk.src[:4], pkt[12:16])
copy(fk.dst[:4], pkt[16:20])
return pkt[:totalLen], true
}
// parseIPv6Prologue is the shared IPv6 tail; the caller has verified len(pkt) >= 40, the version,
// and that the L4 header sits at byte 40.
func (fk *flowKey) parseIPv6Prologue(pkt []byte) ([]byte, bool) {
payloadLen := int(binary.BigEndian.Uint16(pkt[4:6]))
if 40+payloadLen > len(pkt) {
return nil, false
}
fk.isV6 = true
copy(fk.src[:], pkt[8:24])
copy(fk.dst[:], pkt[24:40])
return pkt[:40+payloadLen], true
}
// ipHeadersMatch compares the IP portion of two packet header prefixes for
// byte-for-byte equality on every field that must be identical across coalesced segments.
// Size/IPID/IPCsum are masked out.
// The full DSCP/ECN byte (IPv4 ToS / IPv6 traffic class) is compared, matching Linux kernel GRO:
// segments with differing ECN codepoints must not coalesce,
// otherwise ORing e.g. ECT(0) with ECT(1) would fabricate a false CE (congestion) mark or mark a Not-ECT flow as ECN-capable.
//
// The transport (L4) portion of the header is checked separately by the per-protocol matcher.
func ipHeadersMatch(a, b []byte, isV6 bool) bool {
if isV6 {
// IPv6: [0:4] = version/TC/flow label (TC[1:0] is ECN, so the full TC byte must match),
// [6:40] = next_hdr/hop + src + dst. Skip [4:6] payload_len.
return bytes.Equal(a[:4], b[:4]) && bytes.Equal(a[6:40], b[6:40])
}
// IPv4: [0:2] = version/IHL + DSCP|ECN (full ECN byte must match),
// [6:10] = flags/fragoff/TTL/proto, [12:20] = src+dst.
// Skip [2:4] total len, [4:6] id, [10:12] csum.
return bytes.Equal(a[:2], b[:2]) && bytes.Equal(a[6:10], b[6:10]) && bytes.Equal(a[12:20], b[12:20])
}
// ipv4FlagDF is the Don't Fragment bit in the IPv4 flags byte (header byte 6).
const ipv4FlagDF = 0x40
// ipv4CanCoalesceID reports whether an IPv4 packet whose header starts at
// nextHdr may join a chain whose seed header is seedHdr as segment index seg
// (the seed is segment 0). Kernel GSO re-stamps outgoing segment IDs as
// seed_id+n, so coalescing is only transparent when that re-stamp is either
// harmless (DF set: RFC 6864 atomic datagrams, the ID carries no meaning) or
// reproduces the original IDs exactly (DF clear + IDs already sequential —
// the same admission rule kernel GRO applies). Without this, a DF=0 sender
// with non-sequential IDs (e.g. OpenBSD's randomized IDs) could have IDs
// rewritten into ranges that collide across superpackets, corrupting
// reassembly if the packets are fragmented after the TUN write.
//
// DF itself is guaranteed uniform across a chain by ipHeadersMatch (byte 6
// is inside its compared range), so checking the seed's copy suffices.
func ipv4CanCoalesceID(seedHdr, nextHdr []byte, seg int) bool {
if seedHdr[6]&ipv4FlagDF != 0 {
return true
}
expect := binary.BigEndian.Uint16(seedHdr[4:6]) + uint16(seg)
return binary.BigEndian.Uint16(nextHdr[4:6]) == expect
}
// Arena is an injectable byte-slab that hands out non-overlapping borrowed
// slices via Reserve and releases them in bulk via Reset.
type Arena struct {
buf []byte
}
// NewArena returns an Arena with a pre-allocated backing of the given capacity.
func NewArena(capacity int) *Arena {
return &Arena{buf: make([]byte, 0, capacity)}
}
// Reserve hands out a non-overlapping sz-byte slice from the arena.
// If the request doesn't fit the current backing, a fresh, larger backing is allocated.
// Already-borrowed slices reference the old backing and remain valid until Reset.
func (a *Arena) Reserve(sz int) []byte {
if len(a.buf)+sz > cap(a.buf) {
newCap := max(cap(a.buf)*2, sz)
a.buf = make([]byte, 0, newCap)
}
start := len(a.buf)
a.buf = a.buf[:start+sz]
return a.buf[start : start+sz : start+sz]
}
// Reset releases every slice handed out since the last Reset.
// Callers must not use any previously-borrowed slice after this returns.
// The underlying backing array is retained so subsequent Reserves don't re-allocate.
func (a *Arena) Reset() {
a.buf = a.buf[:0]
}
-112
View File
@@ -1,112 +0,0 @@
package batch
import (
"testing"
"github.com/slackhq/nebula/test"
)
// stagePackets builds the stagedPacket entries Commit would have produced, so dispatch benchmarks
// bypass staging and the sort entirely.
func stagePackets(pkts [][]byte) []stagedPacket {
staged := make([]stagedPacket, len(pkts))
for i, p := range pkts {
pp := testPP(p)
staged[i] = stagedPacket{
pkt: p,
key: SortKey{Epoch: 1, Counter: uint64(i + 1)},
proto: pp.Protocol,
fragAny: pp.FragAny,
ipHdrLen: uint16(pp.IPHdrLen),
}
}
return staged
}
func flushLanes(b *testing.B, m *MultiCoalescer) {
b.Helper()
if m.tcp != nil {
if err := m.tcp.Flush(); err != nil {
b.Fatal(err)
}
}
if m.udp != nil {
if err := m.udp.Flush(); err != nil {
b.Fatal(err)
}
}
if err := m.pt.Flush(); err != nil {
b.Fatal(err)
}
}
// runDispatchBench measures dispatch plus the per-batch lane flush: the post-sort half of the
// batcher, which is where the production profile concentrates.
func runDispatchBench(b *testing.B, pkts [][]byte, batchSize int) {
b.Helper()
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger())
staged := stagePackets(pkts)
b.ReportAllocs()
b.SetBytes(int64(len(pkts[0])))
b.ResetTimer()
for i := 0; i < b.N; i++ {
if err := m.dispatch(staged[i%len(staged)]); err != nil {
b.Fatal(err)
}
if (i+1)%batchSize == 0 {
flushLanes(b, m)
}
}
b.StopTimer()
flushLanes(b, m)
}
// BenchmarkDispatchSingleFlow is the bulk steady state: every packet past the seed appends.
func BenchmarkDispatchSingleFlow(b *testing.B) {
runDispatchBench(b, buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200), tcpCoalesceMaxSegs)
}
// BenchmarkDispatchInterleaved16 stresses the openSlots map: 16 flows round-robined defeats the
// lastSlot cache on every packet.
func BenchmarkDispatchInterleaved16(b *testing.B) {
pkts := buildTCPv4Interleaved(16, tcpCoalesceMaxSegs, 1200)
runDispatchBench(b, pkts, len(pkts))
}
// BenchmarkDispatchAckHeavy alternates MSS data with pure ACKs on one flow — the RX shape of a
// bidirectional transfer (the peer's data and its ACKs of our data share the tunnel direction).
func BenchmarkDispatchAckHeavy(b *testing.B) {
pay := make([]byte, 1200)
var pkts [][]byte
seq := uint32(1000)
for range tcpCoalesceMaxSegs / 2 {
pkts = append(pkts, buildTCPv4(seq, tcpAck, pay))
seq += uint32(len(pay))
pkts = append(pkts, buildTCPv4(seq, tcpAck, nil))
}
runDispatchBench(b, pkts, len(pkts))
}
// BenchmarkDispatchUDPFlow is the QUIC-ish bulk UDP shape.
func BenchmarkDispatchUDPFlow(b *testing.B) {
pay := make([]byte, 1200)
pkts := make([][]byte, udpCoalesceMaxSegs)
for i := range pkts {
pkts[i] = buildUDPv4(2000, 443, pay)
}
runDispatchBench(b, pkts, len(pkts))
}
// BenchmarkDispatchSeedHeavy sets PSH on every packet so each one seeds and immediately closes
// its own slot — the small-write RPC shape, and the upper bound on what the seed path (including
// the parsedTCP-to-slot field transfer) can cost.
func BenchmarkDispatchSeedHeavy(b *testing.B) {
pay := make([]byte, 1200)
pkts := make([][]byte, tcpCoalesceMaxSegs)
seq := uint32(1000)
for i := range pkts {
pkts[i] = buildTCPv4(seq, tcpAckPsh, pay)
seq += uint32(len(pay))
}
runDispatchBench(b, pkts, len(pkts))
}
-76
View File
@@ -1,76 +0,0 @@
package batch
//TODO refactor this away
// This file holds the lanes' self-parsing Commit entries and the proto-checking parsers behind
// them. Production traffic enters the lanes only through MultiCoalescer.dispatch and the At
// parsers; these wrappers reproduce that path (including seal-all on unparseable shapes) on top
// of a local parse, so tests and benches can drive one lane with nothing but a packet.
// parseIPPrologue resolves the IP version, requires the L4 protocol to match wantProto (6 TCP,
// 17 UDP), and defers to the shared per-version cores. Returns the trimmed packet and the L4
// offset; fk must be zero on entry and is filled in place.
func (fk *flowKey) parseIPPrologue(pkt []byte, wantProto byte) ([]byte, int, bool) {
if len(pkt) < 20 {
return nil, 0, false
}
switch pkt[0] >> 4 {
case 4:
if pkt[9] != wantProto {
return nil, 0, false
}
trimmed, ok := fk.parseIPv4Prologue(pkt)
return trimmed, 20, ok
case 6:
if len(pkt) < 40 {
return nil, 0, false
}
if pkt[6] != wantProto {
return nil, 0, false
}
trimmed, ok := fk.parseIPv6Prologue(pkt)
return trimmed, 40, ok
}
return nil, 0, false
}
// parseBase extracts the flow key and IP/TCP offsets for any TCP packet, admissible for
// coalescing or not. Returns false for non-TCP or malformed input.
func (p *parsedTCP) parseBase(pkt []byte) bool {
trimmed, ipHdrLen, ok := p.fk.parseIPPrologue(pkt, ipProtoTCP)
if !ok {
return false
}
return p.parseTail(trimmed, ipHdrLen)
}
// parseBase extracts the flow key and IP/UDP offsets for a UDP packet.
func (p *parsedUDP) parseBase(pkt []byte) bool {
trimmed, ipHdrLen, ok := p.fk.parseIPPrologue(pkt, ipProtoUDP)
if !ok {
return false
}
return p.parseTail(trimmed, ipHdrLen)
}
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
func (c *TCPCoalescer) Commit(pkt []byte) error {
var info parsedTCP
if !info.parseBase(pkt) {
// Unparseable: flow key unknown, seal everything so later data cannot emit ahead of it.
c.sealAllOpen()
c.addVerbatim(pkt)
return nil
}
return c.commitParsed(pkt, &info)
}
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
func (c *UDPCoalescer) Commit(pkt []byte) error {
var info parsedUDP
if !info.parseBase(pkt) {
c.sealAllOpen()
c.addVerbatim(pkt)
return nil
}
return c.commitParsed(pkt, &info)
}
-133
View File
@@ -1,133 +0,0 @@
package batch
import (
"cmp"
"errors"
"io"
"log/slog"
"slices"
"github.com/slackhq/nebula/firewall"
)
// MultiCoalescer stages plaintext packets with their (epoch, counter) sort keys and, at Flush,
// replays them in sender-transmission order into lane-specific batchers selected by L4 protocol.
//
// Sorting before dispatch keeps the ordering story simple: each lane consumes packets in
// transmission order, builds slots in that order, and emits them in creation order. Wire reorder
// inside a flush batch is repaired here, before it can fragment a lane's coalesce chains, so the
// lanes carry no reorder-repair machinery.
//
// The contract is per-tunnel transmission order within each lane, with two exceptions: a pure TCP
// ACK may be overtaken by later same-flow data (it does not close the flow's open chain; a late
// ACK is just a stale ACK), and an unparseable shape seals every open chain in its lane (its flow
// is unknown) and rides the lane as an in-lane verbatim, still in transmission order. Routing
// follows the flow: a flow's non-coalesceable shapes ride its protocol lane rather than falling
// to the later-flushed pt lane.
//
// Cross-lane order (TCP vs UDP vs everything else) is not preserved.
type MultiCoalescer struct {
tcp *TCPCoalescer
udp *UDPCoalescer
pt *Passthrough
// staged holds this batch's packets and sort keys until Flush. Borrowed: the caller keeps
// each pkt alive until Flush returns.
staged []stagedPacket
}
// stagedPacket carries the scalars dispatch needs from the firewall's ParsedPacket, copied by
// value: pp is reused by the caller per packet and must not be retained past Commit.
type stagedPacket struct {
pkt []byte
key SortKey
proto byte
fragAny bool
ipHdrLen uint16
}
// NewMultiCoalescer builds a multi-lane batcher over w, based on available protocol support. The
// staging sort applies even when no GSO lane is available: passthrough-only platforms still get
// transmission-order repair.
func NewMultiCoalescer(w io.Writer, l *slog.Logger) *MultiCoalescer {
m := &MultiCoalescer{
pt: NewPassthrough(w),
staged: make([]stagedPacket, 0, initialSlots),
}
m.tcp = NewTCPCoalescer(w, l)
m.udp = NewUDPCoalescer(w)
return m
}
// Commit stages pkt for the next Flush; dispatch is deferred so it runs on packets already in
// transmission order. key carries the packet's tunnel epoch and message counter. pkt is borrowed:
// the caller must keep it valid until the next Flush and not re-use it, and Flush may patch a
// coalesced packet's headers in place. pp is the firewall's parse of pkt and is borrowed only
// for this call, so the fields dispatch needs are copied here.
func (m *MultiCoalescer) Commit(pkt []byte, key SortKey, pp *firewall.ParsedPacket) error {
m.staged = append(m.staged, stagedPacket{
pkt: pkt,
key: key,
proto: pp.Protocol,
fragAny: pp.FragAny,
ipHdrLen: uint16(pp.IPHdrLen),
})
return nil
}
// compareStaged orders staged packets by (epoch, counter)
func compareStaged(a, b stagedPacket) int {
if c := cmp.Compare(a.key.Epoch, b.key.Epoch); c != 0 {
return c
}
return cmp.Compare(a.key.Counter, b.key.Counter)
}
// dispatch routes one staged packet to its protocol lane (see commitStaged), or to the verbatim
// passthrough when the lane has no GSO support.
func (m *MultiCoalescer) dispatch(sp stagedPacket) error {
switch sp.proto {
case ipProtoTCP:
if m.tcp != nil {
return m.tcp.commitStaged(sp)
}
case ipProtoUDP:
if m.udp != nil {
return m.udp.commitStaged(sp)
}
}
return m.pt.enqueue(sp.pkt)
}
// Flush sorts the staged batch into transmission order, replays it into the lanes, then flushes each lane.
// Drains everything and returns the joined errors; one bad packet does not hold up the rest.
// After Flush returns, committed payload slices may be recycled.
func (m *MultiCoalescer) Flush() error {
// Arrival order is already almost sorted (reorder is the exception), which pdqsort detects
// and handles in near-linear time.
slices.SortFunc(m.staged, compareStaged)
var errs []error
for _, sp := range m.staged {
if err := m.dispatch(sp); err != nil {
errs = append(errs, err)
}
}
clear(m.staged) // drop borrowed pkt refs
m.staged = m.staged[:0]
if m.tcp != nil {
if err := m.tcp.Flush(); err != nil {
errs = append(errs, err)
}
}
if m.udp != nil {
if err := m.udp.Flush(); err != nil {
errs = append(errs, err)
}
}
if err := m.pt.Flush(); err != nil {
errs = append(errs, err)
}
return errors.Join(errs...)
}
-437
View File
@@ -1,437 +0,0 @@
package batch
import (
"bytes"
"encoding/binary"
"io"
"testing"
"github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/test"
)
// keySeq hands out SortKeys with ascending counters in a fixed epoch, for
// tests where commit order IS transmission order.
type keySeq struct {
epoch, counter uint64
}
func (k *keySeq) next() SortKey {
k.counter++
return SortKey{Epoch: k.epoch, Counter: k.counter}
}
// newTestMultiCoalescer builds a batcher over w.
func newTestMultiCoalescer(tb testing.TB, w io.Writer) *MultiCoalescer {
tb.Helper()
return NewMultiCoalescer(w, test.NewLogger())
}
// TestMultiCoalescerRoutesByProto confirms TCP/UDP/other land in the right
// lane: TCP and UDP get coalesced when their lanes are enabled, anything
// else (ICMP here) falls through to plain Write.
func TestMultiCoalescerRoutesByProto(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
m := newTestMultiCoalescer(t, w)
k := &keySeq{epoch: 1}
tcpPay := make([]byte, 1200)
udpPay := make([]byte, 1200)
icmp := make([]byte, 28)
icmp[0] = 0x45
icmp[2] = 0
icmp[3] = 28
icmp[9] = 1
if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(1000, tcpAck, tcpPay))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildTCPv4(2200, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(2200, tcpAck, tcpPay))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil {
t.Fatal(err)
}
if err := m.Commit(icmp, k.next(), testPP(icmp)); err != nil {
t.Fatal(err)
}
if err := m.Flush(); err != nil {
t.Fatal(err)
}
// 1 TCP super (2 segments) + 1 UDP super (2 segments) = 2 gso writes.
if len(w.gsoWrites) != 2 {
t.Fatalf("want 2 gso writes (one TCP + one UDP), got %d", len(w.gsoWrites))
}
if len(w.writes) != 1 {
t.Fatalf("want 1 plain write (ICMP), got %d", len(w.writes))
}
}
// TestMultiCoalescerRestoresTransmissionOrder is the core staging-sort
// property: packets committed out of counter order (wire reorder inside one
// flush batch) are replayed into the lanes in transmission order, so the
// reorder never fragments the coalesce chain — one superpacket, in seq
// order, exactly as if the wire had never reordered. The retransmit shape
// falls out of the same key: a retransmit carries a lower seq but a HIGHER
// counter (it was encrypted later), so it emits after the data it trails.
func TestMultiCoalescerRestoresTransmissionOrder(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
m := newTestMultiCoalescer(t, w)
pay := make([]byte, 1200)
// Transmission order: seq 1000 (c1), 2200 (c2), 3400 (c3).
// Arrival order: 3400, 1000, 2200.
if err := m.Commit(buildTCPv4(3400, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4(3400, tcpAck, pay))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 1, Counter: 1}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildTCPv4(2200, tcpAck, pay), SortKey{Epoch: 1, Counter: 2}, testPP(buildTCPv4(2200, tcpAck, pay))); err != nil {
t.Fatal(err)
}
if err := m.Flush(); err != nil {
t.Fatal(err)
}
if len(w.gsoWrites) != 1 || len(w.writes) != 0 {
t.Fatalf("want 1 gso write (unfragmented chain), got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
}
g := w.gsoWrites[0]
if len(g.pays) != 3 {
t.Fatalf("segs=%d want 3", len(g.pays))
}
const ipHdrLen = 20
if seedSeq := binary.BigEndian.Uint32(g.hdr[ipHdrLen+4 : ipHdrLen+8]); seedSeq != 1000 {
t.Errorf("seed seq=%d want 1000", seedSeq)
}
// Retransmit: seq 1000 again but counter 4 — sorts after seq 4600 (c3).
w.writes, w.gsoWrites, w.order = nil, nil, nil
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 1, Counter: 4}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildTCPv4(4600, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4(4600, tcpAck, pay))); err != nil {
t.Fatal(err)
}
if err := m.Flush(); err != nil {
t.Fatal(err)
}
if len(w.writes) != 2 {
t.Fatalf("want 2 plain writes, got %d (gso=%d)", len(w.writes), len(w.gsoWrites))
}
first := binary.BigEndian.Uint32(w.writes[0][24:28])
second := binary.BigEndian.Uint32(w.writes[1][24:28])
if first != 4600 || second != 1000 {
t.Fatalf("emission (%d, %d), want (4600, 1000): retransmit must not overtake in-flight data", first, second)
}
}
// TestMultiCoalescerRestoresOrderAcrossFlows scrambles two interleaved flows;
// the staging sort must repair each flow into one superpacket without any
// cross-flow contamination.
func TestMultiCoalescerRestoresOrderAcrossFlows(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
m := newTestMultiCoalescer(t, w)
pay := make([]byte, 1200)
// Transmission: A.100 (c1), B.500 (c2), A.1300 (c3), B.1700 (c4).
// Arrival: A.1300, B.1700, A.100, B.500.
if err := m.Commit(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay), SortKey{Epoch: 1, Counter: 4}, testPP(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay), SortKey{Epoch: 1, Counter: 1}, testPP(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay), SortKey{Epoch: 1, Counter: 2}, testPP(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay))); err != nil {
t.Fatal(err)
}
if err := m.Flush(); err != nil {
t.Fatal(err)
}
if len(w.gsoWrites) != 2 {
t.Fatalf("want 2 gso writes (one per flow), got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
}
for i, g := range w.gsoWrites {
if len(g.pays) != 2 {
t.Errorf("gso[%d] segs=%d want 2", i, len(g.pays))
}
const ipHdrLen = 20
seedSeq := binary.BigEndian.Uint32(g.hdr[ipHdrLen+4 : ipHdrLen+8])
sport := binary.BigEndian.Uint16(g.hdr[ipHdrLen : ipHdrLen+2])
switch sport {
case 1000:
if seedSeq != 100 {
t.Errorf("flow A seed seq=%d want 100", seedSeq)
}
case 3000:
if seedSeq != 500 {
t.Errorf("flow B seed seq=%d want 500", seedSeq)
}
default:
t.Errorf("unexpected sport %d", sport)
}
}
}
// TestMultiCoalescerEpochOrdersAcrossRehandshake: a re-handshake replaces
// the tunnel, and the replacement's counter space starts near zero — raw
// counter order would emit the new tunnel's packets first while the old
// tunnel's backlog is still arriving. The epoch key must dominate:
// everything from the old tunnel emits before anything from the new one.
func TestMultiCoalescerEpochOrdersAcrossRehandshake(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
m := newTestMultiCoalescer(t, w)
pay := make([]byte, 1200)
// New session's first data arrives before the old session's last data.
if err := m.Commit(buildTCPv4(2200, tcpAck, pay), SortKey{Epoch: 8, Counter: 1}, testPP(buildTCPv4(2200, tcpAck, pay))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 7, Counter: 9_000_000}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
t.Fatal(err)
}
if err := m.Flush(); err != nil {
t.Fatal(err)
}
// Same flow, contiguous seq, identical headers: after the epoch sort the
// two segments append into one superpacket seeded by the OLD session's
// packet.
if len(w.gsoWrites) != 1 {
t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
}
const ipHdrLen = 20
if seedSeq := binary.BigEndian.Uint32(w.gsoWrites[0].hdr[ipHdrLen+4 : ipHdrLen+8]); seedSeq != 1000 {
t.Errorf("seed seq=%d want 1000 (old session first)", seedSeq)
}
}
// TestMultiCoalescerNoUSOFallsThrough verifies that on a queue without USO
// (older kernel: TSO but no GSO_UDP_L4) the UDP lane never comes up and UDP
// packets still reach the kernel via verbatim rather than being lost.
func TestMultiCoalescerNoUSOFallsThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true, noUSO: true}
m := newTestMultiCoalescer(t, w)
k := &keySeq{epoch: 1}
if m.udp != nil {
t.Fatal("UDP lane must not come up without USO")
}
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv4(1000, 53, make([]byte, 800)))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv4(1000, 53, make([]byte, 800)))); err != nil {
t.Fatal(err)
}
if err := m.Flush(); err != nil {
t.Fatal(err)
}
if len(w.gsoWrites) != 0 {
t.Errorf("UDP must NOT be coalesced when USO disabled, got %d gso writes", len(w.gsoWrites))
}
if len(w.writes) != 2 {
t.Errorf("UDP must pass through as 2 plain writes, got %d", len(w.writes))
}
}
// TestMultiCoalescerNoOffloadsStillSorts covers a queue that can't offload
// anything. Both lane constructors refuse, so every packet rides the
// verbatim lane — but the staging sort still applies, so emission follows
// transmission order even without GSO.
func TestMultiCoalescerNoOffloadsStillSorts(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: false}
m := newTestMultiCoalescer(t, w)
if m.tcp != nil || m.udp != nil {
t.Fatal("no lane may come up without offloads")
}
pkts := [][]byte{
buildTCPv4(1000, tcpAck, make([]byte, 1200)),
buildUDPv4(1000, 53, make([]byte, 800)),
buildTCPv4(2200, tcpAck, make([]byte, 1200)),
}
// Committed in reverse transmission order; keys carry the truth.
for i := len(pkts) - 1; i >= 0; i-- {
if err := m.Commit(pkts[i], SortKey{Epoch: 1, Counter: uint64(i + 1)}, testPP(pkts[i])); err != nil {
t.Fatal(err)
}
}
if err := m.Flush(); err != nil {
t.Fatal(err)
}
if len(w.gsoWrites) != 0 {
t.Errorf("no GSO writes possible, got %d", len(w.gsoWrites))
}
if len(w.writes) != len(pkts) {
t.Fatalf("want %d plain writes, got %d", len(pkts), len(w.writes))
}
// One lane for everything means the sorted order survives end to end.
for i, want := range pkts {
if !bytes.Equal(w.writes[i], want) {
t.Errorf("write %d out of order or corrupt", i)
}
}
}
// buildUDPv6Fragment builds an IPv6 packet whose extension chain is a
// single fragment header (NH=44) naming UDP as the terminal protocol —
// a first fragment (offset 0, MF set) carrying the UDP header and a
// partial payload.
func buildUDPv6Fragment(sport, dport uint16, payload []byte) []byte {
const ipHdrLen = 40
const fragHdrLen = 8
const udpHdrLen = 8
total := ipHdrLen + fragHdrLen + udpHdrLen + len(payload)
pkt := make([]byte, total)
pkt[0] = 0x60
binary.BigEndian.PutUint16(pkt[4:6], uint16(total-ipHdrLen))
pkt[6] = 44 // fragment extension header
pkt[7] = 64
pkt[8] = 0xfe
pkt[9] = 0x80
pkt[23] = 1
pkt[24] = 0xfe
pkt[25] = 0x80
pkt[39] = 2
pkt[40] = ipProtoUDP // fragment's next header
binary.BigEndian.PutUint16(pkt[42:44], 0x0001) // offset 0, MF set
binary.BigEndian.PutUint32(pkt[44:48], 0x1badf00) // identification
binary.BigEndian.PutUint16(pkt[48:50], sport)
binary.BigEndian.PutUint16(pkt[50:52], dport)
binary.BigEndian.PutUint16(pkt[52:54], uint16(udpHdrLen+len(payload)))
copy(pkt[56:], payload)
return pkt
}
// TestMultiCoalescerIPv6FragmentStaysInLane locks in extension-header
// routing: a fragment whose chain terminates in UDP must ride the UDP lane
// as an in-lane verbatim — emitted ahead of later same-flow datagrams —
// not the verbatim lane, which flushes after every coalescer lane and
// would reorder it behind data that arrived after it.
func TestMultiCoalescerIPv6FragmentStaysInLane(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
m := newTestMultiCoalescer(t, w)
k := &keySeq{epoch: 1}
if err := m.Commit(buildUDPv6Fragment(2000, 53, make([]byte, 512)), k.next(), testPP(buildUDPv6Fragment(2000, 53, make([]byte, 512)))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
t.Fatal(err)
}
if err := m.Flush(); err != nil {
t.Fatal(err)
}
if len(w.writes) != 1 {
t.Fatalf("want the fragment as 1 plain write, got %d", len(w.writes))
}
if len(w.gsoWrites) != 1 {
t.Fatalf("want the two whole datagrams coalesced into 1 gso write, got %d", len(w.gsoWrites))
}
// Transmission order was fragment-then-data; same-lane routing must keep it.
if w.order[0] != "write" {
t.Fatalf("fragment must be emitted before later data (in-lane verbatim), order=%v", w.order)
}
}
// TestMultiCoalescerFragmentSealsUDPChains: an unparseable datagram
// (fragment) seals every open UDP chain, so datagrams from before and after
// it land in separate superpackets and the fragment holds its transmission-
// order position between them.
func TestMultiCoalescerFragmentSealsUDPChains(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
m := newTestMultiCoalescer(t, w)
k := &keySeq{epoch: 1}
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildUDPv6Fragment(2000, 53, make([]byte, 512)), k.next(), testPP(buildUDPv6Fragment(2000, 53, make([]byte, 512)))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
t.Fatal(err)
}
if err := m.Flush(); err != nil {
t.Fatal(err)
}
if len(w.gsoWrites) != 2 {
t.Fatalf("want 2 gso writes (chains sealed around the fragment), got %d", len(w.gsoWrites))
}
if len(w.writes) != 1 {
t.Fatalf("want the fragment as 1 plain write, got %d", len(w.writes))
}
want := []string{"gso", "write", "gso"}
if len(w.order) != 3 || w.order[0] != want[0] || w.order[1] != want[1] || w.order[2] != want[2] {
t.Fatalf("emission order = %v, want %v", w.order, want)
}
}
// TestMultiCoalescerNoTSOFallsThrough mirrors the no-TSO case.
func TestMultiCoalescerNoTSOFallsThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true, noTSO: true}
m := newTestMultiCoalescer(t, w)
k := &keySeq{epoch: 1}
if m.tcp != nil {
t.Fatal("TCP lane must not come up without TSO")
}
pay := make([]byte, 1200)
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), k.next(), testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
t.Fatal(err)
}
if err := m.Commit(buildTCPv4(2200, tcpAck, pay), k.next(), testPP(buildTCPv4(2200, tcpAck, pay))); err != nil {
t.Fatal(err)
}
if err := m.Flush(); err != nil {
t.Fatal(err)
}
if len(w.gsoWrites) != 0 {
t.Errorf("TCP must NOT be coalesced when TSO disabled, got %d gso writes", len(w.gsoWrites))
}
if len(w.writes) != 2 {
t.Errorf("TCP must pass through as 2 plain writes, got %d", len(w.writes))
}
}
// testPP derives the ParsedPacket newPacket would produce for the packet
// shapes the tests build: plain v4/v6, v4 with options or fragment bits set,
// and the single-fragment-header v6 shape from buildUDPv6Fragment. Anything
// unrecognizable stays zero (proto 0 routes to the passthrough lane).
func testPP(pkt []byte) *firewall.ParsedPacket {
pp := &firewall.ParsedPacket{}
if len(pkt) < 20 {
return pp
}
switch pkt[0] >> 4 {
case 4:
pp.Protocol = pkt[9]
pp.IPHdrLen = int(pkt[0]&0x0f) * 4
pp.FragAny = binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0
case 6:
pp.Protocol = pkt[6]
pp.IPHdrLen = 40
if pp.Protocol == 44 { // fragment extension header
pp.Protocol = pkt[40]
pp.IPHdrLen = 48
pp.FragAny = true
}
}
return pp
}
-38
View File
@@ -1,38 +0,0 @@
package batch
import (
"io"
)
// Passthrough is MultiCoalescer's verbatim lane: no batching, packets are written at Flush in the
// order enqueued.
type Passthrough struct {
out io.Writer
slots [][]byte
}
func NewPassthrough(w io.Writer) *Passthrough {
return &Passthrough{
out: w,
slots: make([][]byte, 0, 128),
}
}
// enqueue accepts one packet, already sorted into transmission order by dispatch.
func (p *Passthrough) enqueue(pkt []byte) error {
p.slots = append(p.slots, pkt)
return nil
}
func (p *Passthrough) Flush() error {
var firstErr error
for _, s := range p.slots {
_, err := p.out.Write(s)
if err != nil && firstErr == nil {
firstErr = err
}
}
clear(p.slots)
p.slots = p.slots[:0]
return firstErr
}
-472
View File
@@ -1,472 +0,0 @@
package batch
import (
"bytes"
"encoding/binary"
"io"
"log/slog"
"github.com/slackhq/nebula/overlay/tio"
)
// ipProtoTCP is the IANA protocol number for TCP. Defined here to help Windows out.
const ipProtoTCP = 6
// tcpCoalesceBufSize caps total bytes per superpacket. Mirrors the kernel's
// sk_gso_max_size of ~64KiB; anything beyond this would be rejected anyway.
const tcpCoalesceBufSize = 65535
// tcpCoalesceMaxSegs caps how many segments we'll coalesce into a single
// superpacket. Keeping this well below the kernel's TSO ceiling bounds latency.
const tcpCoalesceMaxSegs = 64
// coalesceSlot is one entry in the coalescer's ordered event queue. A verbatim slot holds a single
// borrowed packet emitted as-is (pure ACK, non-admissible TCP, unparseable, or oversize seed); a
// non-verbatim slot is an in-progress coalesced superpacket. payIovs are borrowed slices of the
// caller's plaintext buffers; the caller must keep them alive until Flush.
type coalesceSlot struct {
verbatim bool
// rawPkt is borrowed: the whole packet for verbatim slots, the seed packet for coalesce
// slots. A slot that never grows past one segment is emitted from rawPkt so its original
// (already valid) L4 checksum ships DATA_VALID instead of making the kernel recompute it.
// A multi-segment slot's superpacket header is rawPkt's, patched in place at flush.
rawPkt []byte
fk flowKey
hdrLen int
ipHdrLen int
isV6 bool
gsoSize int
numSeg int
totalPay int
nextSeq uint32
payIovs [][]byte
}
// TCPCoalescer accumulates adjacent in-flow TCP data segments across multiple concurrent flows and
// emits each flow's run as a single TSO superpacket via tio.GSOWriter. Input must be in sender
// transmission order (MultiCoalescer sorts by (epoch, counter) before dispatch); slots are emitted
// in creation order, so emission reproduces transmission order except for the pure-ACK case in
// commitParsed. Owns no locks; one coalescer per TUN write queue.
type TCPCoalescer struct {
w tio.GSOWriter
// slots is the ordered event queue. Flush walks it once and emits each
// entry as either a WriteGSO (coalesced) or a w.Write (verbatim).
slots []*coalesceSlot
// openSlots maps a flow key to its open slot so new segments can extend an in-progress
// superpacket in O(1). Removal is what closes a chain: on PSH or a short last segment, on a
// non-admissible packet for the flow, or in Flush.
openSlots map[flowKey]*coalesceSlot
// lastSlot caches the most recently touched open slot. Bulk traffic
// arrives in same-flow runs (single-flow steady state, or GRO bursts
// under multi-flow), so comparing the incoming key against the cached
// slot's own fk lets the hot path skip the map lookup (and the aeshash
// of a 38-byte key) for the length of each run.
// Kept in lockstep with openSlots: nil whenever the slot it pointed
// at is removed.
lastSlot *coalesceSlot
pool []*coalesceSlot // free list for reuse
l *slog.Logger
}
// NewTCPCoalescer wraps w, returning nil if w can't accept GSO_TCP writes.
func NewTCPCoalescer(w io.Writer, l *slog.Logger) *TCPCoalescer {
gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP)
if !ok {
return nil
}
return &TCPCoalescer{
w: gw,
slots: make([]*coalesceSlot, 0, initialSlots),
openSlots: make(map[flowKey]*coalesceSlot, initialSlots),
pool: make([]*coalesceSlot, 0, initialSlots),
l: l,
}
}
// parsedTCP holds the fields extracted from a single parse so later steps
// (admission, slot lookup, canAppend) don't re-walk the header.
type parsedTCP struct {
fk flowKey
ipHdrLen int
hdrLen int
payLen int
seq uint32
flags byte
}
// parseAt extracts the flow key and IP/TCP offsets for a packet the dispatcher already knows is
// TCP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on
// entry and is filled in place; see flowKey.parseIPAt for why. Returns false for malformed input
// or any shape that must not coalesce (IPv4 options/fragmentation, IPv6 extension headers).
func (p *parsedTCP) parseAt(pkt []byte, ipHdrLen int) bool {
trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen)
if !ok {
return false
}
return p.parseTail(trimmed, ipHdrLen)
}
// parseTail layers the TCP-header parse on a validated IP prologue. pkt is the trimmed packet;
// fk's addresses are already filled.
func (p *parsedTCP) parseTail(pkt []byte, ipHdrLen int) bool {
if len(pkt) < ipHdrLen+20 {
return false
}
tcpOff := int(pkt[ipHdrLen+12]>>4) * 4
if tcpOff < 20 || tcpOff > 60 {
return false
}
if len(pkt) < ipHdrLen+tcpOff {
return false
}
p.ipHdrLen = ipHdrLen
p.hdrLen = ipHdrLen + tcpOff
p.payLen = len(pkt) - p.hdrLen
p.fk.sport = binary.BigEndian.Uint16(pkt[ipHdrLen : ipHdrLen+2])
p.fk.dport = binary.BigEndian.Uint16(pkt[ipHdrLen+2 : ipHdrLen+4])
p.seq = binary.BigEndian.Uint32(pkt[ipHdrLen+4 : ipHdrLen+8])
p.flags = pkt[ipHdrLen+13]
return true
}
// TCP flag bits (byte 13 of the TCP header). Only the bits the coalescer consults are named;
// FIN/SYN/RST/URG/CWR are rejected by the negative mask in commitParsed.
const (
tcpFlagPsh = 0x08
tcpFlagAck = 0x10
tcpFlagEce = 0x40
)
// sealAllOpen closes every open coalesce chain. Called for unparseable packets: the flow key is
// unknown, so any open chain could otherwise absorb later data and emit it ahead of this packet.
func (c *TCPCoalescer) sealAllOpen() {
clear(c.openSlots)
c.lastSlot = nil
}
// sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips
// hashing the 38-byte key when no chains are open (e.g. ack-dominant queues).
func (c *TCPCoalescer) sealFlow(fk flowKey) {
if len(c.openSlots) == 0 {
return
}
if last := c.lastSlot; last != nil && last.fk == fk {
c.lastSlot = nil
}
delete(c.openSlots, fk)
}
// commitStaged commits one staged packet dispatch routed to this lane. A shape the lane cannot
// coalesce (any fragmentation, unparseable header) seals every open chain
// and rides the lane as an in-lane verbatim, still in transmission order.
func (c *TCPCoalescer) commitStaged(sp stagedPacket) error {
if sp.fragAny {
c.sealAllOpen()
c.addVerbatim(sp.pkt)
return nil
}
var info parsedTCP
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
c.sealAllOpen()
c.addVerbatim(sp.pkt)
return nil
}
return c.commitParsed(sp.pkt, &info)
}
// commitParsed commits one parsed TCP packet. The caller (dispatch, via parseAt) supplies a
// valid parse so the header is not re-walked here.
func (c *TCPCoalescer) commitParsed(pkt []byte, info *parsedTCP) error {
// Admission: only ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE may ride a coalesce chain. CWR marks a
// one-shot congestion transition the receiver must observe at a segment boundary. NB: AccECN
// reuses CWR as ACE counter bits; revisit this check if inner hosts adopt AccECN.
if info.flags&tcpFlagAck == 0 || info.flags&^(tcpFlagAck|tcpFlagPsh|tcpFlagEce) != 0 {
// SYN/FIN/RST/URG/CWR must be observed in sequence. Seal the flow's open slot so later
// in-flow packets cannot extend it and emit ahead of this verbatim.
c.sealFlow(info.fk)
c.addVerbatim(pkt)
return nil
}
if info.payLen == 0 {
// Pure ACK: no ordering obligation toward the flow's data. Delivering it after
// later-transmitted data only makes it a stale ACK, which receivers ignore. Not sealing
// keeps a bidirectional flow's data run coalescing across interleaved peer ACKs, matching
// kernel GRO. This is the only place emission deviates from transmission order.
c.addVerbatim(pkt)
return nil
}
// Cached-slot fast path. Arrival isn't per-packet interleaved even with
// many flows: wire-side GRO delivers runs of same-flow packets
// (deliverSegments splits a superdatagram into up to 64), so the cache
// hits for the length of each run and a miss costs one fk compare
// before the map lookup carries the weight.
var open *coalesceSlot
if last := c.lastSlot; last != nil && last.fk == info.fk {
open = last
} else {
open = c.openSlots[info.fk]
}
if open != nil {
if c.canAppend(open, pkt, info) {
if c.appendPayload(open, pkt, info) {
// Chain closed (PSH or short segment): stop extending it.
c.sealFlow(info.fk)
} else {
c.lastSlot = open
}
return nil
}
// Can't extend (seq gap from upstream loss, header change, or a full
// chain): evict it from openSlots and fall through to seed a fresh slot.
c.sealFlow(info.fk)
}
c.seed(pkt, info)
return nil
}
func (c *TCPCoalescer) Flush() error {
var first error
for _, s := range c.slots {
var err error
if s.verbatim || s.numSeg == 1 {
// A slot that never grew is byte-identical to its seed packet; ship the original so
// its valid checksum rides the DATA_VALID path instead of a kernel software csum.
// rawPkt is only mutated once numSeg >= 2 (PSH propagate, flush patches), so it is
// pristine here.
_, err = c.w.Write(s.rawPkt)
} else {
err = c.flushSlot(s)
}
if err != nil && first == nil {
first = err
}
c.release(s)
}
clear(c.slots)
c.slots = c.slots[:0]
clear(c.openSlots)
c.lastSlot = nil
return first
}
func (c *TCPCoalescer) addVerbatim(pkt []byte) {
s := c.take()
s.verbatim = true
s.rawPkt = pkt
c.slots = append(c.slots, s)
}
func (c *TCPCoalescer) seed(pkt []byte, info *parsedTCP) {
if info.hdrLen+info.payLen > tcpCoalesceBufSize {
// Pathological shape that can't ride a superpacket; emit as-is. No chain for this flow can
// be open here (commitParsed evicts before seeding), so sealFlow is defense in depth
// against a stale cache entry absorbing later data.
c.sealFlow(info.fk)
c.addVerbatim(pkt)
return
}
s := c.take()
s.verbatim = false
// rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is
// the superpacket header flushSlot patches in place.
s.rawPkt = pkt
s.hdrLen = info.hdrLen
s.ipHdrLen = info.ipHdrLen
s.isV6 = info.fk.isV6
s.fk = info.fk
s.gsoSize = info.payLen
s.numSeg = 1
s.totalPay = info.payLen
s.nextSeq = info.seq + uint32(info.payLen)
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
c.slots = append(c.slots, s)
if info.flags&tcpFlagPsh == 0 {
c.openSlots[info.fk] = s
c.lastSlot = s
} else {
// PSH on the seed closes the chain immediately; it is never registered as open.
// Drop any stale entry for this flow too (defense in depth, unreachable if lastSlot's lockstep invariant holds).
c.sealFlow(info.fk)
}
}
// canAppend reports whether info's packet extends the slot's seed: same header shape and stable
// contents, adjacent seq, not oversized. A closed chain never reaches here; closing removes the
// slot from openSlots, the only path in. The header fields read from rawPkt are always pristine:
// the only pre-flush mutation is the PSH propagate, which also closes the chain.
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info *parsedTCP) bool {
if info.hdrLen != s.hdrLen {
return false
}
if info.seq != s.nextSeq {
return false
}
if s.numSeg >= tcpCoalesceMaxSegs {
return false
}
if info.payLen > s.gsoSize {
return false
}
if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize {
return false
}
// ECE state must be stable across a burst.
// Receivers expect the flag set on every segment of a CE-echoing window or none.
seedFlags := s.rawPkt[s.ipHdrLen+13]
if (seedFlags^info.flags)&tcpFlagEce != 0 {
return false
}
if !s.isV6 && !ipv4CanCoalesceID(s.rawPkt, pkt, s.numSeg) {
return false
}
if !headersMatch(s.rawPkt[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
return false
}
return true
}
// appendPayload folds info's packet into s and reports whether the chain is now closed: the
// segment was sub-gsoSize (kernel TSO allows only the final segment to be short) or carried PSH.
// The caller must deregister a closed slot from openSlots.
func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info *parsedTCP) bool {
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
s.numSeg++
s.totalPay += info.payLen
s.nextSeq = info.seq + uint32(info.payLen)
if info.flags&tcpFlagPsh != 0 {
// Propagate PSH into the seed header so kernel TSO sets it on the last segment. Mutating
// rawPkt is safe: PSH also closes the chain, so no admission check re-reads this header.
s.rawPkt[s.ipHdrLen+13] |= tcpFlagPsh
}
return info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0
}
func (c *TCPCoalescer) take() *coalesceSlot {
if n := len(c.pool); n > 0 {
s := c.pool[n-1]
c.pool[n-1] = nil
c.pool = c.pool[:n-1]
return s
}
return &coalesceSlot{}
}
func (c *TCPCoalescer) release(s *coalesceSlot) {
clear(s.payIovs)
*s = coalesceSlot{payIovs: s.payIovs[:0]}
c.pool = append(c.pool, s)
}
// flushSlot patches the superpacket header in place in rawPkt (total length, IPv4 header
// checksum, pseudo-header checksum seed) and calls WriteGSO. The slot is released right after,
// so nothing re-reads the patched header. Does not remove the slot from c.slots.
func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
total := s.hdrLen + s.totalPay
l4Len := total - s.ipHdrLen
hdr := s.rawPkt[:s.hdrLen]
if s.isV6 {
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
} else {
binary.BigEndian.PutUint16(hdr[2:4], uint16(total))
hdr[10] = 0
hdr[11] = 0
binary.BigEndian.PutUint16(hdr[10:12], ipv4HdrChecksum(hdr[:s.ipHdrLen]))
}
var psum uint32
if s.isV6 {
psum = pseudoSumIPv6(hdr[8:24], hdr[24:40], ipProtoTCP, l4Len)
} else {
psum = pseudoSumIPv4(hdr[12:16], hdr[16:20], ipProtoTCP, l4Len)
}
tcsum := s.ipHdrLen + 16
binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum))
return c.w.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoTCP)
}
// headersMatch compares two IP+TCP header prefixes for byte-for-byte
// equality on every field that must be identical across coalesced
// segments. Size/IPID/IPCsum/seq/flags/tcpCsum are masked out.
func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
if len(a) != len(b) {
return false
}
if !ipHeadersMatch(a, b, isV6) {
return false
}
// TCP: compare [0:4] ports, [8:13] ack+dataoff, [14:16] window,
// [18:tcpHdrLen] options (incl. urgent).
tcp := ipHdrLen
if !bytes.Equal(a[tcp:tcp+4], b[tcp:tcp+4]) {
return false
}
if !bytes.Equal(a[tcp+8:tcp+13], b[tcp+8:tcp+13]) {
return false
}
if !bytes.Equal(a[tcp+14:tcp+16], b[tcp+14:tcp+16]) {
return false
}
if !bytes.Equal(a[tcp+18:], b[tcp+18:]) {
return false
}
return true
}
// ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must
// already have its checksum field zeroed) and returns the folded/inverted
// 16-bit value to store.
func ipv4HdrChecksum(hdr []byte) uint16 {
var sum uint32
for i := 0; i+1 < len(hdr); i += 2 {
sum += uint32(binary.BigEndian.Uint16(hdr[i : i+2]))
}
if len(hdr)%2 == 1 {
sum += uint32(hdr[len(hdr)-1]) << 8
}
for sum>>16 != 0 {
sum = (sum & 0xffff) + (sum >> 16)
}
return ^uint16(sum)
}
// pseudoSumIPv4 / pseudoSumIPv6 build the L4 pseudo-header partial sum
// expected by the virtio NEEDS_CSUM kernel path: the 32-bit accumulator
// before folding. proto selects the L4 (TCP or UDP); the UDP coalescer
// reuses these helpers.
func pseudoSumIPv4(src, dst []byte, proto byte, l4Len int) uint32 {
var sum uint32
sum += uint32(binary.BigEndian.Uint16(src[0:2]))
sum += uint32(binary.BigEndian.Uint16(src[2:4]))
sum += uint32(binary.BigEndian.Uint16(dst[0:2]))
sum += uint32(binary.BigEndian.Uint16(dst[2:4]))
sum += uint32(proto)
sum += uint32(l4Len)
return sum
}
func pseudoSumIPv6(src, dst []byte, proto byte, l4Len int) uint32 {
var sum uint32
for i := 0; i < 16; i += 2 {
sum += uint32(binary.BigEndian.Uint16(src[i : i+2]))
sum += uint32(binary.BigEndian.Uint16(dst[i : i+2]))
}
sum += uint32(l4Len >> 16)
sum += uint32(l4Len & 0xffff)
sum += uint32(proto)
return sum
}
// foldOnceNoInvert folds the 32-bit accumulator to 16 bits and returns it unchanged (no one's complement).
// This is what virtio NEEDS_CSUM wants in the L4 checksum field
func foldOnceNoInvert(sum uint32) uint16 {
for sum>>16 != 0 {
sum = (sum & 0xffff) + (sum >> 16)
}
return uint16(sum)
}
-214
View File
@@ -1,214 +0,0 @@
package batch
import (
"encoding/binary"
"testing"
"github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/overlay/tio"
"github.com/slackhq/nebula/test"
)
// nopTunWriter is a zero-alloc tio.GSOWriter for benchmarks. Discards
// everything but satisfies the interface the coalescer detects.
type nopTunWriter struct{}
func (nopTunWriter) Write(p []byte) (int, error) { return len(p), nil }
func (nopTunWriter) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, _ tio.GSOProto) error {
return nil
}
func (nopTunWriter) Capabilities() tio.Capabilities {
return tio.Capabilities{TSO: true, USO: true}
}
// buildTCPv4BulkFlow returns a slice of N adjacent ACK-only TCP segments
// on a single 5-tuple, each carrying payloadLen bytes. Seq numbers are
// contiguous so every packet is coalesceable onto the previous one.
func buildTCPv4BulkFlow(n, payloadLen int) [][]byte {
pkts := make([][]byte, n)
pay := make([]byte, payloadLen)
seq := uint32(1000)
for i := range n {
pkts[i] = buildTCPv4(seq, tcpAck, pay)
seq += uint32(payloadLen)
}
return pkts
}
// buildTCPv4Interleaved returns nFlows * perFlow packets with per-flow
// seq continuity but round-robin across flows — worst case for any
// "last-slot" cache.
func buildTCPv4Interleaved(nFlows, perFlow, payloadLen int) [][]byte {
pay := make([]byte, payloadLen)
seqs := make([]uint32, nFlows)
for i := range seqs {
seqs[i] = uint32(1000 + i*1000000)
}
pkts := make([][]byte, 0, nFlows*perFlow)
for range perFlow {
for f := range nFlows {
sport := uint16(10000 + f)
pkts = append(pkts, buildTCPv4Ports(sport, 2000, seqs[f], tcpAck, pay))
seqs[f] += uint32(payloadLen)
}
}
return pkts
}
// buildTCPv4RunInterleaved returns nFlows*perFlow packets delivered in
// runs of runLen per flow — the arrival pattern wire-side GRO actually
// produces (deliverSegments splits each superdatagram into up to 64
// same-flow packets back to back). Contrast with buildTCPv4Interleaved's
// per-packet round-robin, the adversarial worst case for a last-slot cache.
func buildTCPv4RunInterleaved(nFlows, perFlow, runLen, payloadLen int) [][]byte {
pay := make([]byte, payloadLen)
seqs := make([]uint32, nFlows)
for i := range seqs {
seqs[i] = uint32(1000 + i*1000000)
}
pkts := make([][]byte, 0, nFlows*perFlow)
for done := 0; done < perFlow; done += runLen {
for f := range nFlows {
sport := uint16(10000 + f)
for range runLen {
pkts = append(pkts, buildTCPv4Ports(sport, 2000, seqs[f], tcpAck, pay))
seqs[f] += uint32(payloadLen)
}
}
}
return pkts
}
// buildICMPv4 returns a minimal non-TCP packet that takes the verbatim
// branch in Commit.
func buildICMPv4() []byte {
pkt := make([]byte, 28)
pkt[0] = 0x45
binary.BigEndian.PutUint16(pkt[2:4], 28)
pkt[9] = 1 // ICMP
copy(pkt[12:16], []byte{10, 0, 0, 1})
copy(pkt[16:20], []byte{10, 0, 0, 2})
return pkt
}
// runCommitBench drives Commit over pkts batchSize at a time, flushing
// between batches, and reports per-packet cost.
func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
b.Helper()
c := newTestTCPCoalescer(b, nopTunWriter{})
b.ReportAllocs()
b.SetBytes(int64(len(pkts[0])))
b.ResetTimer()
for i := 0; i < b.N; i++ {
pkt := pkts[i%len(pkts)]
if err := c.Commit(pkt); err != nil {
b.Fatal(err)
}
if (i+1)%batchSize == 0 {
if err := c.Flush(); err != nil {
b.Fatal(err)
}
}
}
// Drain any trailing partial batch so slot state doesn't leak across runs.
_ = c.Flush()
}
// BenchmarkCommitSingleFlow is the bulk-TCP steady state: one flow,
// contiguous seq, 1200-byte payloads. Every packet past the seed should
// append onto the open slot. This is the case we most care about.
func BenchmarkCommitSingleFlow(b *testing.B) {
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
runCommitBench(b, pkts, tcpCoalesceMaxSegs)
}
// BenchmarkCommitInterleaved4 has 4 concurrent bulk flows round-robined.
// A single-entry fast-path cache will miss on every packet; an N-way
// cache or map lookup carries the weight.
func BenchmarkCommitInterleaved4(b *testing.B) {
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
runCommitBench(b, pkts, len(pkts))
}
// BenchmarkCommitInterleaved16 stresses the map at higher flow counts.
func BenchmarkCommitInterleaved16(b *testing.B) {
pkts := buildTCPv4Interleaved(16, tcpCoalesceMaxSegs, 1200)
runCommitBench(b, pkts, len(pkts))
}
// BenchmarkCommitRunInterleaved4 is 4 concurrent flows arriving in
// GRO-burst runs of 16 — the realistic multi-flow pattern. A last-slot
// cache hits for the length of each run; the per-packet round-robin
// benches above are its worst case.
func BenchmarkCommitRunInterleaved4(b *testing.B) {
pkts := buildTCPv4RunInterleaved(4, tcpCoalesceMaxSegs, 16, 1200)
runCommitBench(b, pkts, len(pkts))
}
// BenchmarkCommitPassthrough exercises the non-TCP branch: parseBase
// bails early and addVerbatim is the only work.
func BenchmarkCommitPassthrough(b *testing.B) {
pkt := buildICMPv4()
pkts := make([][]byte, 64)
for i := range pkts {
pkts[i] = pkt
}
runCommitBench(b, pkts, 64)
}
// BenchmarkCommitNonCoalesceableTCP sends SYN|ACK packets on one flow.
// Each packet takes the "TCP but not admissible" branch which does a
// map delete + verbatim. Measures the seal-without-slot cost.
func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
pay := make([]byte, 0)
pkts := make([][]byte, 64)
for i := range pkts {
pkts[i] = buildTCPv4(uint32(1000+i), tcpSyn|tcpAck, pay)
}
runCommitBench(b, pkts, 64)
}
// runMultiCommitBench drives MultiCoalescer.Commit with in-order keys, so
// it includes the staging sort's already-sorted fast path plus the
// dispatch-time parse — the full steady-state cost of the batcher. The
// ParsedPackets are precomputed: in production they fall out of the
// firewall's newPacket, which this bench does not model.
func runMultiCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
b.Helper()
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger())
pps := make([]*firewall.ParsedPacket, len(pkts))
for i, p := range pkts {
pps[i] = testPP(p)
}
b.ReportAllocs()
b.SetBytes(int64(len(pkts[0])))
b.ResetTimer()
for i := 0; i < b.N; i++ {
j := i % len(pkts)
if err := m.Commit(pkts[j], SortKey{Epoch: 1, Counter: uint64(i + 1)}, pps[j]); err != nil {
b.Fatal(err)
}
if (i+1)%batchSize == 0 {
if err := m.Flush(); err != nil {
b.Fatal(err)
}
}
}
_ = m.Flush()
}
// BenchmarkMultiCommitSingleFlow is the multi-lane analogue of
// BenchmarkCommitSingleFlow — same workload but routed through the
// dispatcher. The delta vs the single-lane bench measures dispatcher
// overhead.
func BenchmarkMultiCommitSingleFlow(b *testing.B) {
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
runMultiCommitBench(b, pkts, tcpCoalesceMaxSegs)
}
// BenchmarkMultiCommitInterleaved4 mirrors BenchmarkCommitInterleaved4
// through the dispatcher.
func BenchmarkMultiCommitInterleaved4(b *testing.B) {
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
runMultiCommitBench(b, pkts, len(pkts))
}
File diff suppressed because it is too large Load Diff
-59
View File
@@ -1,59 +0,0 @@
package batch
import "net/netip"
const SendBatchCap = 128
// batchWriter is the minimal subset of udp.Conn needed by SendBatch to flush.
type batchWriter interface {
WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error)
}
// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch.
// One SendBatch is owned by each listenIn goroutine; no locking is needed.
// Slots are backed by an Arena (see its docs)
type SendBatch struct {
out batchWriter
bufs [][]byte
dsts []netip.AddrPort
arena *Arena
}
// NewSendBatch makes a SendBatch with batchCap slots and an arenaSize byte buffer for slices to back those slots
func NewSendBatch(out batchWriter, batchCap, arenaSize int) *SendBatch {
return &SendBatch{
out: out,
bufs: make([][]byte, 0, batchCap),
dsts: make([]netip.AddrPort, 0, batchCap),
arena: NewArena(arenaSize),
}
}
func (b *SendBatch) Reserve(sz int) []byte {
return b.arena.Reserve(sz)
}
// Len reports how many packets are queued for the next Flush. Callers use
// it to flush incrementally once a full sendmmsg batch has accumulated,
// bounding how long the first packet of a large read batch waits.
func (b *SendBatch) Len() int { return len(b.bufs) }
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort) {
b.bufs = append(b.bufs, pkt)
b.dsts = append(b.dsts, dst)
}
// Flush writes every queued packet and reports how many actually went out. A short count means some destinations
// were undeliverable; the batch is drained either way.
func (b *SendBatch) Flush() (int, error) {
var err error
written := 0
if len(b.bufs) > 0 {
written, err = b.out.WriteBatch(b.bufs, b.dsts)
}
clear(b.bufs)
b.bufs = b.bufs[:0]
b.dsts = b.dsts[:0]
b.arena.Reset()
return written, err
}
-122
View File
@@ -1,122 +0,0 @@
package batch
import (
"net/netip"
"testing"
)
type fakeBatchWriter struct {
bufs [][]byte
addrs []netip.AddrPort
}
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
// Snapshot — SendBatch.Flush nils its slot pointers right after WriteBatch
// returns, so tests must capture data before that happens.
w.bufs = make([][]byte, len(bufs))
for i, b := range bufs {
cp := make([]byte, len(b))
copy(cp, b)
w.bufs[i] = cp
}
w.addrs = append(w.addrs[:0], addrs...)
return len(bufs), nil
}
func TestSendBatchReserveCommitFlush(t *testing.T) {
fw := &fakeBatchWriter{}
b := NewSendBatch(fw, 4, 32)
ap := netip.MustParseAddrPort("10.0.0.1:4242")
for i := 0; i < 4; i++ {
slot := b.Reserve(32)
if cap(slot) != 32 {
t.Fatalf("slot %d: cap=%d want 32", i, cap(slot))
}
pkt := append(slot[:0], byte(i), byte(i+1), byte(i+2))
b.Commit(pkt, ap)
}
if _, err := b.Flush(); err != nil {
t.Fatalf("Flush: %v", err)
}
if len(fw.bufs) != 4 {
t.Fatalf("WriteBatch got %d bufs want 4", len(fw.bufs))
}
for i, buf := range fw.bufs {
if len(buf) != 3 || buf[0] != byte(i) {
t.Errorf("buf %d: %x", i, buf)
}
if fw.addrs[i] != ap {
t.Errorf("addr %d: got %v want %v", i, fw.addrs[i], ap)
}
}
// Flush again with nothing committed — should be a no-op.
fw.bufs = nil
if _, err := b.Flush(); err != nil {
t.Fatalf("empty Flush: %v", err)
}
if fw.bufs != nil {
t.Fatalf("empty Flush triggered WriteBatch")
}
// Reuse after Flush.
slot := b.Reserve(32)
if cap(slot) != 32 {
t.Fatalf("after Flush Reserve wrong cap: %d", cap(slot))
}
}
func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
fw := &fakeBatchWriter{}
b := NewSendBatch(fw, 3, 8)
ap := netip.MustParseAddrPort("10.0.0.1:80")
for i := 0; i < 3; i++ {
s := b.Reserve(8)
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
b.Commit(pkt, ap)
}
if _, err := b.Flush(); err != nil {
t.Fatalf("Flush: %v", err)
}
for i, buf := range fw.bufs {
if buf[0] != byte(0xA0+i) || buf[1] != byte(0xB0+i) {
t.Errorf("slot %d corrupted: %x", i, buf)
}
}
}
func TestSendBatchGrowPreservesCommitted(t *testing.T) {
fw := &fakeBatchWriter{}
// Tiny initial backing forces a grow on the second Reserve.
b := NewSendBatch(fw, 1, 4)
ap := netip.MustParseAddrPort("10.0.0.1:80")
s1 := b.Reserve(4)
pkt1 := append(s1[:0], 0x11, 0x22, 0x33, 0x44)
b.Commit(pkt1, ap)
s2 := b.Reserve(8) // exceeds remaining cap, triggers grow
pkt2 := append(s2[:0], 0xA, 0xB, 0xC, 0xD, 0xE)
b.Commit(pkt2, ap)
// pkt1 must still be intact even though backing reallocated.
if pkt1[0] != 0x11 || pkt1[3] != 0x44 {
t.Fatalf("first packet corrupted by grow: %x", pkt1)
}
if _, err := b.Flush(); err != nil {
t.Fatalf("Flush: %v", err)
}
if len(fw.bufs) != 2 {
t.Fatalf("got %d bufs want 2", len(fw.bufs))
}
if fw.bufs[0][0] != 0x11 || fw.bufs[0][3] != 0x44 {
t.Errorf("first packet on the wire: %x", fw.bufs[0])
}
if fw.bufs[1][0] != 0xA || fw.bufs[1][4] != 0xE {
t.Errorf("second packet on the wire: %x", fw.bufs[1])
}
}
-345
View File
@@ -1,345 +0,0 @@
package batch
import (
"bytes"
"encoding/binary"
"io"
"github.com/slackhq/nebula/overlay/tio"
)
// ipProtoUDP is the IANA protocol number for UDP.
const ipProtoUDP = 17
// udpCoalesceBufSize caps total bytes per UDP superpacket. Mirrors the
// kernel's gso_max_size; payloads beyond this are emitted as-is.
const udpCoalesceBufSize = 65535
// udpCoalesceMaxSegs caps how many segments we'll coalesce. Kernel UDP-GSO
// accepts up to 64 segments per skb (UDP_MAX_SEGMENTS); stay under that.
const udpCoalesceMaxSegs = 64
// udpSlot is one entry in the UDPCoalescer's ordered event queue.
type udpSlot struct {
verbatim bool
// rawPkt is borrowed: the whole packet for verbatim slots, the seed
// packet for coalesce slots. A coalesce slot that never grows past one
// segment is emitted from rawPkt so its original (already valid) L4
// checksum ships DATA_VALID instead of making the kernel recompute it.
// A multi-segment slot's superpacket header is rawPkt's, patched in place at flush.
rawPkt []byte
fk flowKey
hdrLen int
ipHdrLen int
isV6 bool
gsoSize int // per-segment UDP payload length
numSeg int
totalPay int
payIovs [][]byte
}
// UDPCoalescer accumulates adjacent in-flow UDP datagrams across multiple
// concurrent flows and emits each flow's run as a single GSO_UDP_L4 superpacket via tio.GSOWriter.
// Preserves the in-flow order of packets as they are Commit-ed
//
// Owns no locks; one coalescer per TUN write queue.
type UDPCoalescer struct {
w tio.GSOWriter
slots []*udpSlot
openSlots map[flowKey]*udpSlot
// lastSlot caches the most recently touched open slot; see the
// TCPCoalescer field of the same name. Single-flow QUIC bulk is the
// dominant USO workload, and multi-flow arrival comes in GRO runs, so
// the fk compare beats the map's 38-byte key hash on most packets.
// Kept in lockstep with openSlots: nil whenever the slot it pointed at
// is removed.
lastSlot *udpSlot
pool []*udpSlot
}
func NewUDPCoalescer(w io.Writer) *UDPCoalescer {
gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP)
if !ok {
return nil
}
return &UDPCoalescer{
w: gw,
slots: make([]*udpSlot, 0, initialSlots),
openSlots: make(map[flowKey]*udpSlot, initialSlots),
pool: make([]*udpSlot, 0, initialSlots),
}
}
// parsedUDP holds the fields extracted from a single parse so later steps
// (admission, slot lookup, canAppend) don't re-walk the header.
type parsedUDP struct {
fk flowKey
ipHdrLen int
hdrLen int // ipHdrLen + 8
payLen int
}
// parseAt extracts the flow key and IP/UDP offsets for a packet the dispatcher already knows is
// UDP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on
// entry and is filled in place. Returns false for malformed input or any shape that must not
// coalesce (IPv4 options/fragmentation, IPv6 extension headers).
func (p *parsedUDP) parseAt(pkt []byte, ipHdrLen int) bool {
trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen)
if !ok {
return false
}
return p.parseTail(trimmed, ipHdrLen)
}
// parseTail layers the UDP-header parse on a validated IP prologue. pkt is the trimmed packet;
// fk's addresses are already filled.
func (p *parsedUDP) parseTail(pkt []byte, ipHdrLen int) bool {
if len(pkt) < ipHdrLen+8 {
return false
}
// UDP `length` field: must equal IP-derived length-of-UDP-header-plus-payload.
udpLen := int(binary.BigEndian.Uint16(pkt[ipHdrLen+4 : ipHdrLen+6]))
if udpLen < 8 || udpLen > len(pkt)-ipHdrLen {
return false
}
p.ipHdrLen = ipHdrLen
p.hdrLen = ipHdrLen + 8
p.payLen = udpLen - 8
p.fk.sport = binary.BigEndian.Uint16(pkt[ipHdrLen : ipHdrLen+2])
p.fk.dport = binary.BigEndian.Uint16(pkt[ipHdrLen+2 : ipHdrLen+4])
return true
}
// sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips
// hashing the 38-byte key when no chains are open.
func (c *UDPCoalescer) sealFlow(fk flowKey) {
if len(c.openSlots) == 0 {
return
}
if last := c.lastSlot; last != nil && last.fk == fk {
c.lastSlot = nil
}
delete(c.openSlots, fk)
}
// commitStaged commits one staged packet dispatch routed to this lane. A shape the lane cannot
// coalesce (any fragmentation, unparseable header) seals every open chain — its flow is unknown —
// and rides the lane as an in-lane verbatim, still in transmission order.
func (c *UDPCoalescer) commitStaged(sp stagedPacket) error {
if sp.fragAny {
c.sealAllOpen()
c.addVerbatim(sp.pkt)
return nil
}
var info parsedUDP
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
c.sealAllOpen()
c.addVerbatim(sp.pkt)
return nil
}
return c.commitParsed(sp.pkt, &info)
}
// commitParsed commits one parsed UDP packet. The caller (dispatch, via parseAt) supplies a
// valid parse so the header is not re-walked here.
func (c *UDPCoalescer) commitParsed(pkt []byte, info *parsedUDP) error {
// A zero-length UDP datagram (length == 8) is legal and must reach the TUN, but cannot be
// coalesced.
if info.payLen == 0 {
c.sealFlow(info.fk)
c.addVerbatim(pkt)
return nil
}
// Cached-slot fast path; see the TCPCoalescer equivalent.
var open *udpSlot
if last := c.lastSlot; last != nil && last.fk == info.fk {
open = last
} else {
open = c.openSlots[info.fk]
}
if open != nil {
if c.canAppend(open, pkt, info) {
if c.appendPayload(open, pkt, info) {
// Chain closed (short segment): stop extending it.
c.sealFlow(info.fk)
} else {
c.lastSlot = open
}
return nil
}
// Can't extend: evict it from openSlots and fall through to seed a
// fresh slot.
c.sealFlow(info.fk)
}
c.seed(pkt, info)
return nil
}
func (c *UDPCoalescer) Flush() error {
var first error
for _, s := range c.slots {
var err error
if s.verbatim || s.numSeg == 1 {
// A slot that never grew is byte-identical to the packet it was
// seeded from; ship the original so its valid checksum rides the
// DATA_VALID path instead of paying a kernel software csum.
_, err = c.w.Write(s.rawPkt)
} else {
err = c.flushSlot(s)
}
if err != nil && first == nil {
first = err
}
c.release(s)
}
clear(c.slots)
c.slots = c.slots[:0]
clear(c.openSlots)
c.lastSlot = nil
return first
}
// sealAllOpen closes every open coalesce chain. Called for unparseable packets: the flow key is
// unknown, so any open chain could otherwise absorb later data and emit it ahead of this packet.
func (c *UDPCoalescer) sealAllOpen() {
clear(c.openSlots)
c.lastSlot = nil
}
func (c *UDPCoalescer) addVerbatim(pkt []byte) {
s := c.take()
s.verbatim = true
s.rawPkt = pkt
c.slots = append(c.slots, s)
}
func (c *UDPCoalescer) seed(pkt []byte, info *parsedUDP) {
if info.hdrLen+info.payLen > udpCoalesceBufSize {
// Pathological shape that can't ride a superpacket; emit as-is. No chain for this flow can
// be open here (commitParsed evicts before seeding), so sealFlow is defense in depth
// against a stale cache entry absorbing later data.
c.sealFlow(info.fk)
c.addVerbatim(pkt)
return
}
s := c.take()
s.verbatim = false
// rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is
// the superpacket header flushSlot patches in place.
s.rawPkt = pkt
s.hdrLen = info.hdrLen
s.ipHdrLen = info.ipHdrLen
s.isV6 = info.fk.isV6
s.fk = info.fk
s.gsoSize = info.payLen
s.numSeg = 1
s.totalPay = info.payLen
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
c.slots = append(c.slots, s)
c.openSlots[info.fk] = s
c.lastSlot = s
}
// canAppend reports whether info's packet extends the slot's seed.
// Kernel UDP-GSO requires every segment except possibly the last to be
// exactly gsoSize, and the last may be shorter (≤ gsoSize).
func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info *parsedUDP) bool {
if info.hdrLen != s.hdrLen {
return false
}
if s.numSeg >= udpCoalesceMaxSegs {
return false
}
if info.payLen > s.gsoSize {
return false
}
if s.hdrLen+s.totalPay+info.payLen > udpCoalesceBufSize {
return false
}
// Header reads use rawPkt, which is never mutated before flush. A closed chain never reaches
// here; closing removes the slot from openSlots, the only path in.
if !s.isV6 && !ipv4CanCoalesceID(s.rawPkt, pkt, s.numSeg) {
return false
}
if !udpHeadersMatch(s.rawPkt[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
return false
}
return true
}
// appendPayload folds info's packet into s and reports whether the chain is now closed: kernel
// UDP-GSO requires every segment but the last to be exactly gsoSize, so a short segment must be
// the final one. The caller must deregister a closed slot from openSlots.
func (c *UDPCoalescer) appendPayload(s *udpSlot, pkt []byte, info *parsedUDP) bool {
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
s.numSeg++
s.totalPay += info.payLen
return info.payLen < s.gsoSize
}
func (c *UDPCoalescer) take() *udpSlot {
if n := len(c.pool); n > 0 {
s := c.pool[n-1]
c.pool[n-1] = nil
c.pool = c.pool[:n-1]
return s
}
return &udpSlot{}
}
func (c *UDPCoalescer) release(s *udpSlot) {
// Reset every field, identity ones included; see TCPCoalescer.release.
clear(s.payIovs)
*s = udpSlot{payIovs: s.payIovs[:0]}
c.pool = append(c.pool, s)
}
// flushSlot patches the IP header total length / IPv6 payload length and
// the UDP length to the *total* across all coalesced segments, then seeds
// the UDP checksum field with the pseudo-header partial (single-fold, not
// inverted) per virtio NEEDS_CSUM. The patches land in place in rawPkt; the
// slot is released right after, so nothing re-reads the patched header.
func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
hdr := s.rawPkt[:s.hdrLen]
total := s.hdrLen + s.totalPay // full IP+UDP+all_payloads bytes
l4Len := total - s.ipHdrLen // total UDP (8 + sum of payloads)
if s.isV6 {
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
} else {
binary.BigEndian.PutUint16(hdr[2:4], uint16(total))
hdr[10] = 0
hdr[11] = 0
binary.BigEndian.PutUint16(hdr[10:12], ipv4HdrChecksum(hdr[:s.ipHdrLen]))
}
// UDP length field (offset 4 inside the UDP header) = total UDP size.
binary.BigEndian.PutUint16(hdr[s.ipHdrLen+4:s.ipHdrLen+6], uint16(l4Len))
var psum uint32
if s.isV6 {
psum = pseudoSumIPv6(hdr[8:24], hdr[24:40], ipProtoUDP, l4Len)
} else {
psum = pseudoSumIPv4(hdr[12:16], hdr[16:20], ipProtoUDP, l4Len)
}
udpCsumOff := s.ipHdrLen + 6
binary.BigEndian.PutUint16(hdr[udpCsumOff:udpCsumOff+2], foldOnceNoInvert(psum))
return c.w.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoUDP)
}
// udpHeadersMatch compares two IP+UDP header prefixes for byte-equality on
// every field that must be identical across coalesced segments
func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
if len(a) != len(b) {
return false
}
if !ipHeadersMatch(a, b, isV6) {
return false
}
// UDP: compare sport+dport ([0:4]). Skip length [4:6] and checksum [6:8]:
// length varies (we rewrite at flush) and the checksum will be redone.
udp := ipHdrLen
return bytes.Equal(a[udp:udp+4], b[udp:udp+4])
}
-72
View File
@@ -1,72 +0,0 @@
package batch
import (
"testing"
)
// buildUDPv4BulkFlow returns n equal-size datagrams on one flow — the
// steady state for single-flow QUIC bulk, the workload USO exists for.
func buildUDPv4BulkFlow(n, payloadLen int) [][]byte {
pay := make([]byte, payloadLen)
pkts := make([][]byte, n)
for i := range pkts {
pkts[i] = buildUDPv4(40000, 443, pay)
}
return pkts
}
// buildUDPv4RunInterleaved mirrors buildTCPv4RunInterleaved: nFlows*perFlow
// datagrams arriving in GRO-burst runs of runLen per flow.
func buildUDPv4RunInterleaved(nFlows, perFlow, runLen, payloadLen int) [][]byte {
pay := make([]byte, payloadLen)
pkts := make([][]byte, 0, nFlows*perFlow)
for done := 0; done < perFlow; done += runLen {
for f := range nFlows {
sport := uint16(40000 + f)
for range runLen {
pkts = append(pkts, buildUDPv4(sport, 443, pay))
}
}
}
return pkts
}
// runUDPCommitBench drives UDPCoalescer.Commit over pkts batchSize at a
// time, flushing between batches, and reports per-packet cost.
func runUDPCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
b.Helper()
c := newTestUDPCoalescer(b, nopTunWriter{})
b.ReportAllocs()
b.SetBytes(int64(len(pkts[0])))
b.ResetTimer()
for i := 0; i < b.N; i++ {
pkt := pkts[i%len(pkts)]
if err := c.Commit(pkt); err != nil {
b.Fatal(err)
}
if (i+1)%batchSize == 0 {
if err := c.Flush(); err != nil {
b.Fatal(err)
}
}
}
_ = c.Flush()
}
// BenchmarkUDPCommitSingleFlow is the single-flow bulk steady state.
func BenchmarkUDPCommitSingleFlow(b *testing.B) {
pkts := buildUDPv4BulkFlow(udpCoalesceMaxSegs, 1200)
runUDPCommitBench(b, pkts, udpCoalesceMaxSegs)
}
// BenchmarkUDPCommitInterleaved4 is the adversarial per-packet round-robin.
func BenchmarkUDPCommitInterleaved4(b *testing.B) {
pkts := buildUDPv4RunInterleaved(4, udpCoalesceMaxSegs, 1, 1200)
runUDPCommitBench(b, pkts, len(pkts))
}
// BenchmarkUDPCommitRunInterleaved4 is 4 flows in GRO-burst runs of 16.
func BenchmarkUDPCommitRunInterleaved4(b *testing.B) {
pkts := buildUDPv4RunInterleaved(4, udpCoalesceMaxSegs, 16, 1200)
runUDPCommitBench(b, pkts, len(pkts))
}
-536
View File
@@ -1,536 +0,0 @@
package batch
import (
"bytes"
"encoding/binary"
"io"
"testing"
)
// buildUDPv4 builds a minimal IPv4+UDP packet with the given payload and ports.
func buildUDPv4(sport, dport uint16, payload []byte) []byte {
const ipHdrLen = 20
const udpHdrLen = 8
total := ipHdrLen + udpHdrLen + len(payload)
pkt := make([]byte, total)
pkt[0] = 0x45
pkt[1] = 0x00
binary.BigEndian.PutUint16(pkt[2:4], uint16(total))
binary.BigEndian.PutUint16(pkt[4:6], 0)
binary.BigEndian.PutUint16(pkt[6:8], 0x4000)
pkt[8] = 64
pkt[9] = ipProtoUDP
copy(pkt[12:16], []byte{10, 0, 0, 1})
copy(pkt[16:20], []byte{10, 0, 0, 2})
binary.BigEndian.PutUint16(pkt[20:22], sport)
binary.BigEndian.PutUint16(pkt[22:24], dport)
binary.BigEndian.PutUint16(pkt[24:26], uint16(udpHdrLen+len(payload)))
binary.BigEndian.PutUint16(pkt[26:28], 0)
copy(pkt[28:], payload)
return pkt
}
// buildUDPv6 builds a minimal IPv6+UDP packet.
func buildUDPv6(sport, dport uint16, payload []byte) []byte {
const ipHdrLen = 40
const udpHdrLen = 8
total := ipHdrLen + udpHdrLen + len(payload)
pkt := make([]byte, total)
pkt[0] = 0x60
binary.BigEndian.PutUint16(pkt[4:6], uint16(udpHdrLen+len(payload)))
pkt[6] = ipProtoUDP
pkt[7] = 64
pkt[8] = 0xfe
pkt[9] = 0x80
pkt[23] = 1
pkt[24] = 0xfe
pkt[25] = 0x80
pkt[39] = 2
binary.BigEndian.PutUint16(pkt[40:42], sport)
binary.BigEndian.PutUint16(pkt[42:44], dport)
binary.BigEndian.PutUint16(pkt[44:46], uint16(udpHdrLen+len(payload)))
binary.BigEndian.PutUint16(pkt[46:48], 0)
copy(pkt[48:], payload)
return pkt
}
// newTestUDPCoalescer builds a coalescer over w and fails the test if w can't
// do USO. See newTestTCPCoalescer.
func newTestUDPCoalescer(tb testing.TB, w io.Writer) *UDPCoalescer {
tb.Helper()
c := NewUDPCoalescer(w)
if c == nil {
tb.Fatal("NewUDPCoalescer: writer does not support USO")
}
return c
}
// TestNewUDPCoalescerRefusesWhenGSOUnavailable mirrors the TCP precondition:
// no USO, no coalescer.
func TestNewUDPCoalescerRefusesWhenGSOUnavailable(t *testing.T) {
if c := NewUDPCoalescer(&fakeTunWriter{gsoEnabled: false}); c != nil {
t.Fatalf("want nil for a non-USO writer, got %v", c)
}
if c := NewUDPCoalescer(&plainOnlyWriter{}); c != nil {
t.Fatalf("want nil for a plain writer, got %v", c)
}
}
func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
// ICMP packet
pkt := make([]byte, 28)
pkt[0] = 0x45
binary.BigEndian.PutUint16(pkt[2:4], 28)
pkt[9] = 1
copy(pkt[12:16], []byte{10, 0, 0, 1})
copy(pkt[16:20], []byte{10, 0, 0, 2})
if err := c.Commit(pkt); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
t.Fatalf("ICMP must pass through unchanged: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
}
}
func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pkt := buildUDPv4(1000, 53, make([]byte, 800))
if err := c.Commit(pkt); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
// A slot that never grew past one datagram flushes as a plain Write of
// the original packet bytes: the original (already valid) checksum
// ships via the DATA_VALID path, so the kernel does no csum work.
// WriteGSO is reserved for slots that actually coalesced (>=2 segs).
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
t.Fatalf("single-seg flush: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
}
if !bytes.Equal(w.writes[0], pkt) {
t.Errorf("plain write not byte-identical to committed packet: got %d bytes want %d", len(w.writes[0]), len(pkt))
}
}
func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pay := make([]byte, 1200)
for i := 0; i < 3; i++ {
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
t.Fatal(err)
}
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(w.gsoWrites) != 1 {
t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
}
g := w.gsoWrites[0]
if g.gsoSize != 1200 {
t.Errorf("gsoSize=%d want 1200", g.gsoSize)
}
if len(g.pays) != 3 {
t.Errorf("pay count=%d want 3", len(g.pays))
}
if g.csumStart != 20 {
t.Errorf("csumStart=%d want 20", g.csumStart)
}
// IP totalLen and UDP length must be the TOTAL across all segments —
// the kernel's ip_rcv_core trims skbs to iph->tot_len, so a per-segment
// value would silently drop everything but the first segment. Total =
// IP(20) + UDP(8) + 3*1200 = 3628.
gotTotalLen := binary.BigEndian.Uint16(g.hdr[2:4])
if gotTotalLen != 3628 {
t.Errorf("ipv4 total_len=%d want 3628 (must be total across segments)", gotTotalLen)
}
gotUDPLen := binary.BigEndian.Uint16(g.hdr[20+4 : 20+6])
if gotUDPLen != 8+3*1200 {
t.Errorf("udp len=%d want %d", gotUDPLen, 8+3*1200)
}
}
// Last segment may be shorter, sealing the chain.
func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
full := make([]byte, 1200)
tail := make([]byte, 600)
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
t.Fatal(err)
}
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
t.Fatal(err)
}
if err := c.Commit(buildUDPv4(1000, 53, tail)); err != nil {
t.Fatal(err)
}
// A 4th packet, even same-sized, must NOT join — chain is sealed.
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
// The sealed 3-datagram chain is a real superpacket; the re-seed stays
// single-segment and flushes as a plain write of the original packet.
if len(w.gsoWrites) != 1 || len(w.writes) != 1 {
t.Fatalf("want 1 gso (sealed) + 1 plain (new seed), got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
}
if len(w.gsoWrites[0].pays) != 3 {
t.Errorf("super: want 3 pays, got %d", len(w.gsoWrites[0].pays))
}
if got, want := len(w.writes[0]), 20+8+1200; got != want {
t.Errorf("re-seed plain write len=%d want %d", got, want)
}
}
// A larger-than-gsoSize packet cannot extend the slot — it reseeds.
func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
t.Fatal(err)
}
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 1200))); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
// Both seeds stay single-segment → two plain writes in arrival order.
if len(w.writes) != 2 || len(w.gsoWrites) != 0 {
t.Fatalf("want 2 separate plain writes, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
}
for i, want := range []int{20 + 8 + 800, 20 + 8 + 1200} {
if len(w.writes[i]) != want {
t.Errorf("write %d len=%d want %d", i, len(w.writes[i]), want)
}
}
}
// Different 5-tuples must not coalesce.
func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pay := make([]byte, 800)
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
t.Fatal(err)
}
if err := c.Commit(buildUDPv4(2000, 53, pay)); err != nil {
t.Fatal(err)
}
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
t.Fatal(err)
}
if err := c.Commit(buildUDPv4(2000, 53, pay)); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
// Two flows × 2 datagrams each = 2 superpackets of 2 segments.
if len(w.gsoWrites) != 2 {
t.Fatalf("want 2 gso writes (one per flow), got %d", len(w.gsoWrites))
}
for i, g := range w.gsoWrites {
if len(g.pays) != 2 {
t.Errorf("super %d: want 2 pays, got %d", i, len(g.pays))
}
}
}
// Caps at udpCoalesceMaxSegs.
func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pay := make([]byte, 100)
for i := 0; i < udpCoalesceMaxSegs+5; i++ {
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
t.Fatal(err)
}
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
// First superpacket holds udpCoalesceMaxSegs segments; the spillover
// reseeds a new one.
if len(w.gsoWrites) != 2 {
t.Fatalf("want 2 gso writes (cap then reseed), got %d", len(w.gsoWrites))
}
if len(w.gsoWrites[0].pays) != udpCoalesceMaxSegs {
t.Errorf("first super: pays=%d want %d", len(w.gsoWrites[0].pays), udpCoalesceMaxSegs)
}
if len(w.gsoWrites[1].pays) != 5 {
t.Errorf("second super: pays=%d want 5", len(w.gsoWrites[1].pays))
}
}
// Differing IP ECN codepoints must not coalesce: udpHeadersMatch compares
// the full ToS byte (matching kernel GRO). A CE-marked datagram mid-run
// seals the Not-ECT chain and reseeds; the trailing Not-ECT datagram
// reseeds again. All three stay single-segment, so each ships as a plain
// write of its original bytes, keeping its own codepoint.
func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pay := make([]byte, 800)
pkt0 := buildUDPv4(1000, 53, pay) // ECN=00 (Not-ECT)
pkt1 := buildUDPv4(1000, 53, pay)
pkt1[1] = 0x03 // CE
pkt2 := buildUDPv4(1000, 53, pay) // ECN=00 again
for _, p := range [][]byte{pkt0, pkt1, pkt2} {
if err := c.Commit(p); err != nil {
t.Fatal(err)
}
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(w.writes) != 3 || len(w.gsoWrites) != 0 {
t.Fatalf("want 3 separate plain writes (differing ECN), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
}
wantECN := []byte{0x00, 0x03, 0x00}
for i, p := range w.writes {
if got := p[1] & 0x03; got != wantECN[i] {
t.Errorf("write %d ECN=%#x want %#x", i, got, wantECN[i])
}
}
}
// IPv6 path: same flow, equal-sized → coalesced.
func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pay := make([]byte, 1200)
for i := 0; i < 3; i++ {
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
t.Fatal(err)
}
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(w.gsoWrites) != 1 {
t.Fatalf("want 1 gso write, got %d", len(w.gsoWrites))
}
g := w.gsoWrites[0]
if !g.isV6 {
t.Errorf("expected v6 write")
}
if g.csumStart != 40 {
t.Errorf("csumStart=%d want 40", g.csumStart)
}
// IPv6 payload_len and UDP length must be TOTAL — kernel's
// ip6_rcv_core trims to payload_len + ipv6 hdr size. Total UDP = 8 +
// 3*1200 = 3608.
gotPlen := binary.BigEndian.Uint16(g.hdr[4:6])
if gotPlen != 8+3*1200 {
t.Errorf("ipv6 payload_len=%d want %d (must be total)", gotPlen, 8+3*1200)
}
gotUDPLen := binary.BigEndian.Uint16(g.hdr[40+4 : 40+6])
if gotUDPLen != 8+3*1200 {
t.Errorf("udp len=%d want %d", gotUDPLen, 8+3*1200)
}
}
// DSCP differences must reseed: udpHeadersMatch compares the full ToS byte.
func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pay := make([]byte, 800)
pkt0 := buildUDPv4(1000, 53, pay)
pkt1 := buildUDPv4(1000, 53, pay)
pkt1[1] = 0xb8 // EF DSCP, ECN=0
if err := c.Commit(pkt0); err != nil {
t.Fatal(err)
}
if err := c.Commit(pkt1); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
// Both seeds stay single-segment → two plain writes, no gso.
if len(w.writes) != 2 || len(w.gsoWrites) != 0 {
t.Fatalf("want 2 separate plain writes (different DSCP), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
}
}
// Fragmented IPv4 must not be coalesced.
func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pkt := buildUDPv4(1000, 53, make([]byte, 200))
binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1
if err := c.Commit(pkt); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
t.Fatalf("frag must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
}
}
// A zero-length UDP datagram (UDP length == 8, no payload) is legal and
// must be delivered as a plain single datagram — never coalesced. Seeding
// it into a GSO slot stores an empty payload iovec that panics WriteGSO
// (index-out-of-range on &pay[0]); this is a remote DoS if we ever let it
// reach the GSO path. Regression: must not panic and must be written.
func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pkt := buildUDPv4(1000, 53, nil) // UDP length 8, zero payload
if err := c.Commit(pkt); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
t.Fatalf("zero-length UDP must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
}
if len(w.writes[0]) != len(pkt) {
t.Errorf("delivered %d bytes, want the whole %d-byte datagram", len(w.writes[0]), len(pkt))
}
}
// IPv6 zero-length UDP datagram: same verbatim contract as v4.
func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pkt := buildUDPv6(1000, 53, nil) // UDP length 8, zero payload
if err := c.Commit(pkt); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
t.Fatalf("zero-length IPv6 UDP must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
}
if len(w.writes[0]) != len(pkt) {
t.Errorf("delivered %d bytes, want the whole %d-byte datagram", len(w.writes[0]), len(pkt))
}
}
// A zero-length datagram arriving mid-flow must seal the open chain so the
// datagram after it seeds a fresh superpacket *after* the empty one on the
// wire — per-flow arrival order (full, empty, full) must be preserved.
func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
full := make([]byte, 800)
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
t.Fatal(err)
}
if err := c.Commit(buildUDPv4(1000, 53, nil)); err != nil { // zero-length
t.Fatal(err)
}
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
// The empty datagram sealed the first slot, so the trailing full packet
// can't join it. All three emit as plain writes (the two full datagrams
// stayed single-segment; the empty one is verbatim) in per-flow
// arrival order: full, empty, full.
if len(w.writes) != 3 || len(w.gsoWrites) != 0 {
t.Fatalf("want 3 plain writes, got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
}
for i, want := range []int{20 + 8 + 800, 20 + 8, 20 + 8 + 800} {
if len(w.writes[i]) != want {
t.Errorf("write %d len=%d want %d (order full, empty, full)", i, len(w.writes[i]), want)
}
}
}
// IPv4 with options is not admissible (we require IHL=5).
func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pkt := buildUDPv4(1000, 53, make([]byte, 200))
pkt[0] = 0x46 // IHL = 6 (24-byte IPv4 header — has options)
if err := c.Commit(pkt); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
t.Fatalf("ipv4-with-options must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
}
}
// TestUDPCoalescerNonAtomicSequentialIDsCoalesce mirrors the TCP rule: DF
// clear is fine as long as the IDs already run seed+1 per datagram, so
// kernel USO's re-stamp reproduces them.
func TestUDPCoalescerNonAtomicSequentialIDsCoalesce(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pay := make([]byte, 1200)
for i := range 2 {
pkt := buildUDPv4(40000, 443, pay)
setIPv4ID(pkt, uint16(40+i), false)
if err := c.Commit(pkt); err != nil {
t.Fatal(err)
}
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(w.gsoWrites) != 1 || len(w.gsoWrites[0].pays) != 2 {
t.Fatalf("sequential-ID DF=0 datagrams must coalesce: gso=%d", len(w.gsoWrites))
}
}
// TestUDPCoalescerNonAtomicIDGapReseeds: an ID jump on a DF=0 flow breaks
// the chain; each datagram stays a single-segment slot and flushes as a
// plain write that keeps its own (meaningful) ID.
func TestUDPCoalescerNonAtomicIDGapReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pay := make([]byte, 1200)
p1 := buildUDPv4(40000, 443, pay)
setIPv4ID(p1, 40, false)
p2 := buildUDPv4(40000, 443, pay)
setIPv4ID(p2, 50, false)
if err := c.Commit(p1); err != nil {
t.Fatal(err)
}
if err := c.Commit(p2); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(w.writes) != 2 || len(w.gsoWrites) != 0 {
t.Fatalf("ID gap on DF=0 must reseed: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
}
for i, want := range []uint16{40, 50} {
if id := binary.BigEndian.Uint16(w.writes[i][4:6]); id != want {
t.Errorf("write %d: ID=%d want %d (must be preserved)", i, id, want)
}
}
}
-23
View File
@@ -1,23 +0,0 @@
package checksum
import (
"golang.org/x/sys/cpu"
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
)
//go:noescape
func checksumAVX2(buf []byte, initial uint16) uint16
var hasAVX2 = cpu.X86.HasAVX2
// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with
// initial. It is a drop-in replacement for gvisor's checksum.Checksum that
// dispatches to a hand-written AVX2 routine on amd64 CPUs that support it,
// falling back to gvisor's pure-Go implementation otherwise. The result
// matches gvisor's bit-for-bit for any buffer length and initial seed.
func Checksum(buf []byte, initial uint16) uint16 {
if hasAVX2 {
return checksumAVX2(buf, initial)
}
return gvisorchecksum.Checksum(buf, initial)
}
-157
View File
@@ -1,157 +0,0 @@
#include "textflag.h"
// func checksumAVX2(buf []byte, initial uint16) uint16
//
// Computes the RFC 1071 ones-complement sum of buf, seeded with initial.
//
// Algorithm: sum the buffer treating it as a stream of uint32s in machine
// (little-endian) byte order, accumulating into 64-bit lanes (top 32 bits
// hold cross-add carries at 1 byte / lane / iter we have 32 bits of
// headroom which is far more than the 16 KB/64 KB max practical inputs).
// At the end we fold to 16 bits and byte-swap once to recover the on-wire
// (big-endian) result. RFC 1071 §1.2.B byte-order independence makes this
// equivalent to summing as 16-bit big-endian words.
//
// The ymm accumulators (Y4..Y7) hold 4 uint64 lanes each = 16 parallel
// partial sums. The main loop loads 64 bytes per iter as four 16-byte
// chunks, zero-extending each chunk's four uint32s into a ymm via
// VPMOVZXDQ-from-memory, then VPADDQ into a separate accumulator per
// chunk to break the dep chain. After the vector loop the lane sums are
// horizontally reduced and merged with a scalar accumulator that handles
// the trailing 0..63 bytes plus the (byte-swapped) initial seed.
TEXT ·checksumAVX2(SB), NOSPLIT, $0-34
MOVQ buf_base+0(FP), SI
MOVQ buf_len+8(FP), CX
MOVWQZX initial+24(FP), AX
// Pre-byteswap initial into the LE-summing space so it merges directly
// with the rest of the accumulator. The final fold's bswap16 will undo
// this and convert the whole result back to BE.
XCHGB AH, AL
CMPQ CX, $32
JLT scalar_tail
VPXOR Y4, Y4, Y4
VPXOR Y5, Y5, Y5
VPXOR Y6, Y6, Y6
VPXOR Y7, Y7, Y7
CMPQ CX, $64
JLT loop32
loop64:
VPMOVZXDQ (SI), Y0
VPMOVZXDQ 16(SI), Y1
VPMOVZXDQ 32(SI), Y2
VPMOVZXDQ 48(SI), Y3
VPADDQ Y0, Y4, Y4
VPADDQ Y1, Y5, Y5
VPADDQ Y2, Y6, Y6
VPADDQ Y3, Y7, Y7
ADDQ $64, SI
SUBQ $64, CX
CMPQ CX, $64
JGE loop64
loop32:
CMPQ CX, $32
JLT reduce_vec
VPMOVZXDQ (SI), Y0
VPMOVZXDQ 16(SI), Y1
VPADDQ Y0, Y4, Y4
VPADDQ Y1, Y5, Y5
ADDQ $32, SI
SUBQ $32, CX
JMP loop32
reduce_vec:
// Combine the four ymm accumulators into Y4.
VPADDQ Y5, Y4, Y4
VPADDQ Y7, Y6, Y6
VPADDQ Y6, Y4, Y4
// Horizontally reduce Y4's four uint64 lanes to a single scalar.
VEXTRACTI128 $1, Y4, X5
VPADDQ X5, X4, X4
VPSHUFD $0x4e, X4, X5
VPADDQ X5, X4, X4
VMOVQ X4, R8
VZEROUPPER
ADDQ R8, AX
ADCQ $0, AX
scalar_tail:
// Handle remaining 0..63 bytes (or the entire buffer if it was < 32).
CMPQ CX, $8
JLT tail4
loop8:
ADDQ (SI), AX
ADCQ $0, AX
ADDQ $8, SI
SUBQ $8, CX
CMPQ CX, $8
JGE loop8
tail4:
CMPQ CX, $4
JLT tail2
MOVL (SI), R8
ADDQ R8, AX
ADCQ $0, AX
ADDQ $4, SI
SUBQ $4, CX
tail2:
CMPQ CX, $2
JLT tail1
MOVWQZX (SI), R8
ADDQ R8, AX
ADCQ $0, AX
ADDQ $2, SI
SUBQ $2, CX
tail1:
TESTQ CX, CX
JZ fold
MOVBQZX (SI), R8
ADDQ R8, AX
ADCQ $0, AX
fold:
// Fold the 64-bit accumulator to 16 bits via four rounds, mirroring
// gvisor's reduce(). Each pair (split, add) halves the live width;
// the truncation steps absorb the single bit that may be left over
// after each add so the next round's bound holds.
// 64 33 bits.
MOVQ AX, R8
SHRQ $32, R8
MOVL AX, AX
ADDQ R8, AX
// 33 32 bits. AX += (AX>>32); truncate to 32. AX is now ≤ 0xFFFF_FFFF.
MOVQ AX, R8
SHRQ $32, R8
ADDQ R8, AX
MOVL AX, AX
// 32 17 bits.
MOVQ AX, R8
SHRQ $16, R8
MOVWQZX AX, AX
ADDQ R8, AX
// 17 16 bits. AX += (AX>>16); the trailing MOVW truncates bit 16.
MOVQ AX, R8
SHRQ $16, R8
ADDQ R8, AX
// AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip
// to big-endian to match the gvisor API contract.
XCHGB AH, AL
MOVW AX, ret+32(FP)
RET
-12
View File
@@ -1,12 +0,0 @@
package checksum
//go:noescape
func checksumNEON(buf []byte, initial uint16) uint16
// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with
// initial. It is a drop-in replacement for gvisor's checksum.Checksum
// that dispatches to a hand-written NEON routine. NEON is mandatory in
// armv8 so no feature check is needed.
func Checksum(buf []byte, initial uint16) uint16 {
return checksumNEON(buf, initial)
}
-143
View File
@@ -1,143 +0,0 @@
#include "textflag.h"
// func checksumNEON(buf []byte, initial uint16) uint16
//
// Mirrors the algorithm in checksum_amd64.s: sum the buffer treating it as
// a stream of uint32s in machine (little-endian) byte order, accumulating
// into 64-bit lanes that have ample carry headroom; fold and byte-swap once
// at the very end to recover the on-wire (big-endian) result.
//
// Each loop iteration loads 64 bytes via VLD1.P into V0..V3 (4 Q regs).
// VUADDW takes the low two uint32 lanes of a Q reg, zero-extends them to
// uint64, and adds them into a 2×uint64 accumulator; VUADDW2 does the same
// for the high two lanes. Four ymm-equivalent accumulators (V8..V11) get
// updated twice per iter to break the dep chain. Tail bytes go through a
// scalar ADCS chain seeded with the byte-swapped initial.
TEXT ·checksumNEON(SB), NOSPLIT, $0-34
MOVD buf_base+0(FP), R0
MOVD buf_len+8(FP), R1
MOVHU initial+24(FP), R2
// Pre-byteswap initial into the LE-summing space so it merges directly
// with the rest of the accumulator.
REV16W R2, R2
MOVD ZR, R3 // scalar accumulator
CMP $32, R1
BLT scalar_tail
VEOR V8.B16, V8.B16, V8.B16
VEOR V9.B16, V9.B16, V9.B16
VEOR V10.B16, V10.B16, V10.B16
VEOR V11.B16, V11.B16, V11.B16
CMP $64, R1
BLT loop16_init
loop64:
VLD1.P 64(R0), [V0.B16, V1.B16, V2.B16, V3.B16]
VUADDW V0.S2, V8.D2, V8.D2
VUADDW2 V0.S4, V9.D2, V9.D2
VUADDW V1.S2, V10.D2, V10.D2
VUADDW2 V1.S4, V11.D2, V11.D2
VUADDW V2.S2, V8.D2, V8.D2
VUADDW2 V2.S4, V9.D2, V9.D2
VUADDW V3.S2, V10.D2, V10.D2
VUADDW2 V3.S4, V11.D2, V11.D2
SUB $64, R1, R1
CMP $64, R1
BGE loop64
loop16_init:
CMP $16, R1
BLT reduce_vec
loop16:
VLD1.P 16(R0), [V0.B16]
VUADDW V0.S2, V8.D2, V8.D2
VUADDW2 V0.S4, V9.D2, V9.D2
SUB $16, R1, R1
CMP $16, R1
BGE loop16
reduce_vec:
// Combine the four accumulators into V8.
VADD V9.D2, V8.D2, V8.D2
VADD V11.D2, V10.D2, V10.D2
VADD V10.D2, V8.D2, V8.D2
// Horizontal-add the two lanes of V8.D2 into a single uint64.
VADDP V8.D2, V8.D2, V8.D2
VMOV V8.D[0], R8
ADDS R8, R3, R3
ADC ZR, R3, R3
scalar_tail:
CMP $8, R1
BLT tail4
loop8:
MOVD.P 8(R0), R8
ADDS R8, R3, R3
ADC ZR, R3, R3
SUB $8, R1, R1
CMP $8, R1
BGE loop8
tail4:
CMP $4, R1
BLT tail2
MOVWU.P 4(R0), R8
ADDS R8, R3, R3
ADC ZR, R3, R3
SUB $4, R1, R1
tail2:
CMP $2, R1
BLT tail1
MOVHU.P 2(R0), R8
ADDS R8, R3, R3
ADC ZR, R3, R3
SUB $2, R1, R1
tail1:
CBZ R1, fold
MOVBU (R0), R8
ADDS R8, R3, R3
ADC ZR, R3, R3
fold:
// Merge the byte-swapped initial into our LE-form accumulator.
ADDS R2, R3, R3
ADC ZR, R3, R3
// 64 33 bits.
LSR $32, R3, R8
AND $0xffffffff, R3, R3
ADD R8, R3, R3
// 33 32 (truncate after adding bit 32 back).
LSR $32, R3, R8
ADD R8, R3, R3
AND $0xffffffff, R3, R3
// 32 17.
LSR $16, R3, R8
AND $0xffff, R3, R3
ADD R8, R3, R3
// 17 16 (truncation absorbs bit 16 below).
LSR $16, R3, R8
ADD R8, R3, R3
// AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip
// to big-endian to match the gvisor API contract. REV16W swaps bytes
// within each 16-bit halfword of the low 32 bits, so it acts as a
// 16-bit byte-swap on the live low 16.
REV16W R3, R3
AND $0xffff, R3, R3
MOVH R3, ret+32(FP)
RET
-10
View File
@@ -1,10 +0,0 @@
//go:build !amd64 && !arm64
package checksum
import gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
// Checksum delegates to gvisor on architectures without a hand-written body.
func Checksum(buf []byte, initial uint16) uint16 {
return gvisorchecksum.Checksum(buf, initial)
}
-232
View File
@@ -1,232 +0,0 @@
package checksum
import (
"fmt"
"math/rand/v2"
"testing"
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
)
// archImpl names one checksum function under test. The per-arch
// export_*_test.go files enumerate the hand-written implementations so the
// suite compares each one against gvisor directly, regardless of which one
// the public Checksum dispatches to on the running CPU. Testing only the
// dispatcher was tautological wherever it resolved to the gvisor fallback
// (non-AVX2 amd64, fallback architectures) — gvisor compared with itself,
// assembly untested, suite green.
type archImpl struct {
name string
fn func([]byte, uint16) uint16
available bool
}
// implsUnderTest is the public dispatcher plus every arch implementation.
func implsUnderTest() []archImpl {
return append([]archImpl{{name: "dispatch", fn: Checksum, available: true}}, archImpls...)
}
// requireAvailable skips loudly when the running CPU can't execute an
// implementation — visible in test output, unlike the old silent tautology.
func requireAvailable(t *testing.T, impl archImpl) {
t.Helper()
if !impl.available {
t.Skipf("%s not supported on this CPU; its assembly is NOT tested in this run", impl.name)
}
}
// TestChecksumMatchesGvisor walks lengths from 0 to 4096, with several initial
// seeds and a handful of starting alignments, asserting that each local
// implementation matches gvisor's reference bit-for-bit.
func TestChecksumMatchesGvisor(t *testing.T) {
for _, impl := range implsUnderTest() {
t.Run(impl.name, func(t *testing.T) {
requireAvailable(t, impl)
rng := rand.New(rand.NewPCG(1, 2))
const padFront = 16
// Random pool large enough for the longest case + alignment slop.
pool := make([]byte, 4096+padFront)
for i := range pool {
pool[i] = byte(rng.Uint32())
}
seeds := []uint16{0, 0x0001, 0xabcd, 0xffff, 0x1234, 0xfedc}
offsets := []int{0, 1, 2, 3, 4, 5, 7, 8, 15, 16}
for length := 0; length <= 4096; length++ {
for _, seed := range seeds {
for _, off := range offsets {
if off+length > len(pool) {
continue
}
buf := pool[off : off+length]
want := gvisorchecksum.Checksum(buf, seed)
got := impl.fn(buf, seed)
if got != want {
t.Fatalf("len=%d off=%d seed=%#x: got %#04x want %#04x",
length, off, seed, got, want)
}
}
}
}
})
}
}
// TestChecksumPatternedBuffers exercises specific byte patterns that have
// historically tripped up checksum implementations: all-zero, all-0xff,
// alternating, and ascending sequences.
func TestChecksumPatternedBuffers(t *testing.T) {
for _, impl := range implsUnderTest() {
t.Run(impl.name, func(t *testing.T) {
requireAvailable(t, impl)
for length := 0; length <= 256; length++ {
patterns := map[string][]byte{
"zeros": make([]byte, length),
"ones": bytes(length, 0xff),
"alternating": pattern(length, []byte{0xa5, 0x5a}),
"ascending": ascending(length),
}
for name, buf := range patterns {
for _, seed := range []uint16{0, 0xffff, 0x8000} {
want := gvisorchecksum.Checksum(buf, seed)
got := impl.fn(buf, seed)
if got != want {
t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x",
name, length, seed, got, want)
}
}
}
}
})
}
}
func bytes(n int, v byte) []byte {
b := make([]byte, n)
for i := range b {
b[i] = v
}
return b
}
func pattern(n int, p []byte) []byte {
b := make([]byte, n)
for i := range b {
b[i] = p[i%len(p)]
}
return b
}
func ascending(n int) []byte {
b := make([]byte, n)
for i := range b {
b[i] = byte(i)
}
return b
}
// TestChecksumTailPaths targets every combination of (SIMD body iterations,
// trailing tail bytes) the asm handlers walk through. The tail handlers
// peel off 8 → 4 → 2 → 1 byte chunks in turn; this test exercises each by
// constructing lengths of the form 64*k + tail for tail ∈ [0, 63] and a
// representative spread of k values, including k=0 (no main loop, all tail)
// and k=1 (one main loop iter, then tail). It's explicit coverage for
// payload sizes that are odd, not divisible by 4, by 8, or by 32.
func TestChecksumTailPaths(t *testing.T) {
for _, impl := range implsUnderTest() {
t.Run(impl.name, func(t *testing.T) {
requireAvailable(t, impl)
rng := rand.New(rand.NewPCG(42, 17))
const padFront = 16
const maxK = 8
pool := make([]byte, 64*maxK+padFront+64)
for i := range pool {
pool[i] = byte(rng.Uint32())
}
seeds := []uint16{0, 0xffff, 0xabcd}
offsets := []int{0, 1, 3, 7, 15} // mix of aligned and odd starts
for k := 0; k <= maxK; k++ {
for tail := 0; tail < 64; tail++ {
length := 64*k + tail
for _, seed := range seeds {
for _, off := range offsets {
if off+length > len(pool) {
continue
}
buf := pool[off : off+length]
want := gvisorchecksum.Checksum(buf, seed)
got := impl.fn(buf, seed)
if got != want {
t.Fatalf("k=%d tail=%d (len=%d) off=%d seed=%#x: got %#04x want %#04x",
k, tail, length, off, seed, got, want)
}
}
}
}
}
})
}
}
// BenchmarkChecksumTailSizes covers payload sizes that aren't clean multiples
// of the SIMD body's 32-byte (amd64) or 16-byte (arm64) chunks, so the tail
// handler is meaningfully on the hot path. Sizes are picked to either exercise
// every tail branch (tiny lengths) or sit slightly off realistic packet
// boundaries (e.g. 1499 = MTU 1).
func BenchmarkChecksumTailSizes(b *testing.B) {
sizes := []int{
1, 3, 7, 15, 31, // sub-SIMD; entire work is scalar tail
33, 35, 47, 63, // one loop32 + assorted tails
65, 95, 127, // one loop64 + assorted tails
1447, 1471, 1499, 1501, // around MTU
8191, 8193, // around USO
65531, 65533, // near the kernel max
}
for _, size := range sizes {
buf := make([]byte, size)
for i := range buf {
buf[i] = byte(i)
}
b.Run(fmt.Sprintf("size=%d/local", size), func(b *testing.B) {
b.SetBytes(int64(size))
for i := 0; i < b.N; i++ {
_ = Checksum(buf, 0)
}
})
b.Run(fmt.Sprintf("size=%d/gvisor", size), func(b *testing.B) {
b.SetBytes(int64(size))
for i := 0; i < b.N; i++ {
_ = gvisorchecksum.Checksum(buf, 0)
}
})
}
}
// BenchmarkChecksum compares the local Checksum to gvisor's at sizes that
// match real traffic: a TCP/IP header (60), a typical MSS (1448), a typical
// USO size (8192), and the kernel's max GSO superpacket (65535).
func BenchmarkChecksum(b *testing.B) {
for _, size := range []int{60, 1448, 8192, 65535} {
buf := make([]byte, size)
for i := range buf {
buf[i] = byte(i)
}
b.Run(fmt.Sprintf("size=%d/local", size), func(b *testing.B) {
b.SetBytes(int64(size))
for i := 0; i < b.N; i++ {
_ = Checksum(buf, 0)
}
})
b.Run(fmt.Sprintf("size=%d/gvisor", size), func(b *testing.B) {
b.SetBytes(int64(size))
for i := 0; i < b.N; i++ {
_ = gvisorchecksum.Checksum(buf, 0)
}
})
}
}
-11
View File
@@ -1,11 +0,0 @@
package checksum
// archImpls exposes every hand-written implementation on this architecture
// so the tests exercise them directly, independent of what the public
// Checksum dispatches to on the running CPU. Without this, running the
// suite on a non-AVX2 machine compared gvisor against itself and left the
// assembly untested — silently. available=false makes the test skip loudly
// instead.
var archImpls = []archImpl{
{name: "avx2", fn: checksumAVX2, available: hasAVX2},
}
-8
View File
@@ -1,8 +0,0 @@
package checksum
// archImpls exposes every hand-written implementation on this architecture
// for direct testing; see export_amd64_test.go for the rationale. NEON is
// mandatory in armv8, so it is always available.
var archImpls = []archImpl{
{name: "neon", fn: checksumNEON, available: true},
}
-7
View File
@@ -1,7 +0,0 @@
//go:build !amd64 && !arm64
package checksum
// No hand-written implementations on this architecture; the dispatcher is
// pure gvisor and there is nothing separate to test.
var archImpls []archImpl
+3 -13
View File
@@ -4,25 +4,15 @@ import (
"io"
"net/netip"
"github.com/slackhq/nebula/overlay/tio"
"github.com/slackhq/nebula/routing"
)
// defaultBatchBufSize is the per-Queue scratch size for Read on backends
// that don't do TSO segmentation. 65535 covers any single IP packet.
const defaultBatchBufSize = 65535
type Device interface {
io.Closer
io.ReadWriteCloser
Activate() error
Networks() []netip.Prefix
Name() string
RoutesFor(netip.Addr) routing.Gateways
// Queues returns the device's packet queues, opening additional ones as
// needed until there are n. Platforms without multiqueue support return
// their single queue regardless of n, so callers must size reader loops
// to len(result), not n; implementations never return more than n. An
// error means a queue that should have opened could not; the caller owns
// cleanup via Close. Called once, during interface activation.
Queues(n int) ([]tio.Queue, error)
SupportsMultiqueue() bool
NewMultiQueueReader() (io.ReadWriteCloser, error)
}
+10 -5
View File
@@ -3,9 +3,10 @@
package overlaytest
import (
"errors"
"io"
"net/netip"
"github.com/slackhq/nebula/overlay/tio"
"github.com/slackhq/nebula/routing"
)
@@ -30,16 +31,20 @@ func (NoopTun) Name() string {
return "noop"
}
func (NoopTun) Read() ([]tio.Packet, error) {
return nil, nil
func (NoopTun) Read([]byte) (int, error) {
return 0, nil
}
func (NoopTun) Write([]byte) (int, error) {
return 0, nil
}
func (NoopTun) Queues(int) ([]tio.Queue, error) {
return []tio.Queue{NoopTun{}}, nil
func (NoopTun) SupportsMultiqueue() bool {
return false
}
func (NoopTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
return nil, errors.New("unsupported")
}
func (NoopTun) Close() error {
-44
View File
@@ -1,44 +0,0 @@
//go:build linux && !android
// +build linux,!android
package tio
import (
"os"
"golang.org/x/sys/unix"
)
// blockOn parks the calling goroutine until fd is ready or shutdownFd signals teardown.
// (events is POLLIN for reads, POLLOUT for writes)
// It builds the pollfd array on the stack every call, so concurrent callers on the same Queue never share Revents storage.
//
// Returns os.ErrClosed when shutdown was signaled (POLLIN on shutdownFd)
// or either fd reported a problem condition (POLLHUP|POLLNVAL|POLLERR).
func blockOn(fd, shutdownFd int32, events int16) error {
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
pfds := [2]unix.PollFd{
{Fd: fd, Events: events},
{Fd: shutdownFd, Events: unix.POLLIN},
}
var err error
for {
_, err = unix.Poll(pfds[:], -1)
if err != unix.EINTR {
break
}
}
tunEvents := pfds[0].Revents
shutdownEvents := pfds[1].Revents
// Check err before trusting the potentially bogus bits we just got.
if err != nil {
return err
}
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
return os.ErrClosed
}
if tunEvents&problemFlags != 0 {
return os.ErrClosed
}
return nil
}
-101
View File
@@ -1,101 +0,0 @@
//go:build linux && !android
// +build linux,!android
package tio
import (
"encoding/binary"
"errors"
"fmt"
"log/slog"
"sync/atomic"
"golang.org/x/sys/unix"
)
type offloadQueueSet struct {
pq []*Offload
// pqi is exactly the same as pq, but stored as the interface type
pqi []Queue
shutdownFd int
// usoEnabled is true when newTun successfully negotiated TUN_F_USO4|6 with the kernel.
// Queues created by Add inherit this and surface it via Offload.USOSupported so coalescers can gate USO emission.
usoEnabled bool
closed atomic.Bool
// l is handed to each queue for its bad-vnet-header drop logging.
l *slog.Logger
}
// NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do TSO segmentation.
// usoEnabled tells downstream queues whether the kernel agreed to deliver/accept GSO_UDP_L4 superpackets.
func NewOffloadQueueSet(usoEnabled bool, l *slog.Logger) (QueueSet, error) {
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
if err != nil {
return nil, fmt.Errorf("failed to create eventfd: %w", err)
}
out := &offloadQueueSet{
pq: []*Offload{},
pqi: []Queue{},
shutdownFd: shutdownFd,
usoEnabled: usoEnabled,
l: l,
}
return out, nil
}
func (c *offloadQueueSet) Queues() []Queue {
return c.pqi
}
func (c *offloadQueueSet) Add(fd int) error {
if c.closed.Load() {
return errors.New("queue set already closed")
}
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled, c.l)
if err != nil {
return err
}
c.pq = append(c.pq, x)
c.pqi = append(c.pqi, x)
return nil
}
func (c *offloadQueueSet) wakeForShutdown() error {
var buf [8]byte
binary.NativeEndian.PutUint64(buf[:], 1)
_, err := unix.Write(c.shutdownFd, buf[:])
return err
}
func (c *offloadQueueSet) Close() error {
if c.closed.Swap(true) {
return nil
}
errs := []error{}
// Signal all readers blocked in poll to wake up and exit.
// They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
if err := c.wakeForShutdown(); err != nil {
errs = append(errs, err)
}
// Close the per-queue tun fds; this also unblocks any in-flight reads.
for _, x := range c.pq {
if err := x.Close(); err != nil {
errs = append(errs, err)
}
}
// Close the shutdown eventfd last: every reader's pollfd set references it,
// so it must outlive the wake + per-queue teardown above.
if err := unix.Close(c.shutdownFd); err != nil {
errs = append(errs, err)
}
c.shutdownFd = -1
return errors.Join(errs...)
}
-91
View File
@@ -1,91 +0,0 @@
//go:build linux && !android
// +build linux,!android
package tio
import (
"encoding/binary"
"errors"
"fmt"
"sync/atomic"
"golang.org/x/sys/unix"
)
type pollQueueSet struct {
pq []*Poll
// pqi is exactly the same as pq, but stored as the interface type
pqi []Queue
shutdownFd int
closed atomic.Bool
}
func NewPollQueueSet() (QueueSet, error) {
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
if err != nil {
return nil, fmt.Errorf("failed to create eventfd: %w", err)
}
out := &pollQueueSet{
pq: []*Poll{},
pqi: []Queue{},
shutdownFd: shutdownFd,
}
return out, nil
}
func (c *pollQueueSet) Queues() []Queue {
return c.pqi
}
func (c *pollQueueSet) Add(fd int) error {
if c.closed.Load() {
return errors.New("queue set already closed")
}
x, err := newPoll(fd, c.shutdownFd)
if err != nil {
return err
}
c.pq = append(c.pq, x)
c.pqi = append(c.pqi, x)
return nil
}
func (c *pollQueueSet) wakeForShutdown() error {
var buf [8]byte
binary.NativeEndian.PutUint64(buf[:], 1)
_, err := unix.Write(int(c.shutdownFd), buf[:])
return err
}
func (c *pollQueueSet) Close() error {
if c.closed.Swap(true) {
return nil
}
errs := []error{}
// Signal all readers blocked in poll to wake up and exit.
// They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
if err := c.wakeForShutdown(); err != nil {
errs = append(errs, err)
}
// Close the per-queue tun fds; this also unblocks any in-flight reads.
for _, x := range c.pq {
if err := x.Close(); err != nil {
errs = append(errs, err)
}
}
// Close the shutdown eventfd last: every reader's pollfd set references it,
// so it must outlive the wake + per-queue teardown above.
if err := unix.Close(c.shutdownFd); err != nil {
errs = append(errs, err)
}
c.shutdownFd = -1
return errors.Join(errs...)
}
-65
View File
@@ -1,65 +0,0 @@
//go:build linux && !android && !e2e_testing
package tio
import "testing"
// fakeBatch stands in for batch.TxBatcher inside the bench — same shape
// of pointer-capturing closure that sendInsideMessage builds.
type fakeBatch struct{ buf [65536]byte }
func (b *fakeBatch) Reserve(sz int) []byte { return b.buf[:sz] }
func (b *fakeBatch) Commit([]byte) {}
type fakeHostInfo struct {
remoteIndexId uint32
counter uint64
}
type fakeIface struct {
rebindCount uint8
hi *fakeHostInfo
}
// BenchmarkSegmentSuperpacketAllocsTSO measures allocation per
// SegmentSuperpacket call when a closure captures pointer-bearing
// receivers — the realistic shape of sendInsideMessage's closure.
func BenchmarkSegmentSuperpacketAllocsTSO(b *testing.B) {
const mss = 1400
const numSeg = 32
pkt := buildTSOv6(mss*numSeg, mss)
gso := GSOInfo{
Size: mss,
HdrLen: 60, // 40 (IPv6) + 20 (TCP)
CsumStart: 40,
Proto: GSOProtoTCP,
}
p := Packet{Bytes: pkt, GSO: gso}
hi := &fakeHostInfo{remoteIndexId: 0xdeadbeef}
f := &fakeIface{rebindCount: 7, hi: hi}
fb := &fakeBatch{}
// SegmentSuperpacket consumes pkt destructively; refresh from a master
// copy each iter (matches the production pattern where every TUN read
// hands the segmenter a fresh kernel-supplied buffer).
master := append([]byte(nil), pkt...)
work := make([]byte, len(pkt))
p.Bytes = work
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
copy(work, master)
err := SegmentSuperpacket(p, func(seg []byte) error {
out := fb.Reserve(16 + len(seg) + 16)
out[0] = byte(f.rebindCount)
out[1] = byte(hi.counter)
hi.counter++
fb.Commit(out)
return nil
})
if err != nil {
b.Fatalf("SegmentSuperpacket: %v", err)
}
}
}
-16
View File
@@ -1,16 +0,0 @@
//go:build !linux || android
package tio
import "fmt"
func protoFromGSOType(_ uint8) (GSOProto, error) {
return 0, fmt.Errorf("GSO unsupported")
}
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
if pkt.GSO.IsSuperpacket() {
return fmt.Errorf("tio: GSO superpacket on platform without segmentation support")
}
return fn(pkt.Bytes)
}
-49
View File
@@ -1,49 +0,0 @@
package tio
import "io"
// singleQueue adapts a legacy one-datagram-per-Read source into a Queue.
// Read fills a private scratch buffer and returns exactly one Packet whose
// Bytes borrow from that buffer, valid only until the next Read, per the Queue contract.
// Single-reader like every Queue; Write is exactly as safe for concurrent use as the underlying source's Write.
type singleQueue struct {
rw io.ReadWriter
closer io.Closer // nil: Close is a no-op (the source is shared and owned elsewhere)
buf []byte
ret [1]Packet
}
// NewSingleQueue wraps a one-datagram-per-Read ReadWriteCloser (a legacy tun device) into a Queue.
// bufSize is the per-queue read scratch size and must be at least the largest datagram the source can return.
// Close closes rwc.
func NewSingleQueue(rwc io.ReadWriteCloser, bufSize int) Queue {
return &singleQueue{rw: rwc, closer: rwc, buf: make([]byte, bufSize)}
}
// NewSingleQueueNoClose is NewSingleQueue for a source owned by someone else,
// e.g. several queues sharing one device. Close on the returned Queue is a
// no-op so one queue can't tear the shared source out from under its
// siblings; the owner remains responsible for closing the source itself.
func NewSingleQueueNoClose(rw io.ReadWriter, bufSize int) Queue {
return &singleQueue{rw: rw, buf: make([]byte, bufSize)}
}
func (q *singleQueue) Read() ([]Packet, error) {
n, err := q.rw.Read(q.buf)
if err != nil {
return nil, err
}
q.ret[0] = Packet{Bytes: q.buf[:n]}
return q.ret[:], nil
}
func (q *singleQueue) Write(p []byte) (int, error) {
return q.rw.Write(p)
}
func (q *singleQueue) Close() error {
if q.closer == nil {
return nil
}
return q.closer.Close()
}
-147
View File
@@ -1,147 +0,0 @@
package tio
import (
"io"
)
// QueueSet holds one or many Queue objects and helps close them in an orderly way.
type QueueSet interface {
io.Closer
Queues() []Queue
// Add takes a tun fd, adds it to the set, and prepares it for use as a Queue.
Add(fd int) error
}
// Capabilities advertises which kernel offload features a Queue successfully negotiated.
// Callers consult this to decide which coalescers to wire onto the write path.
type Capabilities struct {
// TSO means the FD was opened with IFF_VNET_HDR and the kernel agreed to TUN_F_TSO4|TSO6,
// and WriteGSO with GSOProtoTCP is safe.
TSO bool
// USO means the kernel additionally agreed to TUN_F_USO4|USO6,
// so WriteGSO with GSOProtoUDP is safe. Linux ≥ 6.2.
USO bool
}
// Queue is a readable/writable Poll queue.
// Concurrency contract: a single read goroutine drives Read; plain Write is safe for concurrent callers;
// WriteGSO (on Queues that implement GSOWriter) is single-writer per queue.
//
// Close on an individual Queue does NOT unblock a Read parked in poll — closing an fd
// never wakes its pollers. Orderly teardown goes through the owning QueueSet's Close,
// which first signals a shared shutdown eventfd every reader polls alongside its own fd.
// That eventfd is a set-wide kill switch: once signaled, every Queue in the set returns
// os.ErrClosed from Read, so it cannot be used to stop a single Queue.
type Queue interface {
io.Closer
// Read returns one or more packets.
// The returned Packet.Bytes slices are borrowed from the Queue's internal buffer and are only valid
// until the next Read or Close on this Queue.
// A Packet may carry a GSO/USO superpacket (see GSOInfo)
// Single-reader only: not safe for concurrent Reads (it reuses per-queue rx scratch each call).
Read() ([]Packet, error)
// Write emits a single packet on the plaintext (outside→inside) delivery path.
// Safe for concurrent use.
Write(p []byte) (int, error)
}
// Packet is the unit Queue.Read returns.
// Bytes points into the queue's internal buffer and is only valid until the next Read or Close on the queue that produced it.
// GSO is the zero value for an already-segmented IP datagram;
// when non-zero it describes a kernel-supplied TSO/USO superpacket the caller must segment before consuming.
type Packet struct {
Bytes []byte
GSO GSOInfo
}
// GSOInfo describes a kernel-supplied superpacket sitting in Packet.Bytes.
// The zero value means Bytes is one regular IP datagram and no segmentation is required.
type GSOInfo struct {
// Size is the GSO segment size: max payload bytes per segment
// (== TCP MSS for TSO, == UDP payload chunk for USO). Zero means not a superpacket.
Size uint16
// HdrLen is the total L3+L4 header length within Bytes (already corrected via correctHdrLen, so safe to slice on).
HdrLen uint16
// CsumStart is the L4 header offset inside Bytes (== L3 header length).
CsumStart uint16
// Proto picks the L4 protocol (TCP or UDP) so the segmenter knows which checksum/header layout to apply.
Proto GSOProto
}
// IsSuperpacket reports whether g describes a multi-segment GSO/USO
// superpacket that needs segmentation before its bytes can be encrypted and sent on the wire.
func (g GSOInfo) IsSuperpacket() bool { return g.Size > 0 }
// Clone returns a Packet whose Bytes is a freshly allocated copy of p.Bytes,
// safe to retain past the next Read or Close on the originating Queue.
// GSO metadata is copied verbatim.
// Use this only when a caller needs the data to outlive the borrowed-slice contract.
func (p Packet) Clone() Packet {
if p.Bytes == nil {
return p
}
cp := make([]byte, len(p.Bytes))
copy(cp, p.Bytes)
return Packet{Bytes: cp, GSO: p.GSO}
}
// CapsProvider is an optional interface implemented by Queues that negotiate kernel offload features at open time.
// Callers pick a write-path coalescer based on the result.
// Queues that don't implement it are treated as having no offload capability.
type CapsProvider interface {
Capabilities() Capabilities
}
// GSOProto selects the L4 protocol for a GSO superpacket.
// Determines which VIRTIO_NET_HDR_GSO_* type the writer stamps and which checksum offset
// inside the transport header virtio NEEDS_CSUM expects.
type GSOProto uint8
const (
GSOProtoUnknown GSOProto = iota
GSOProtoTCP
GSOProtoUDP
)
// GSOWriter is implemented by Queues that can emit a TCP or UDP superpacket
// assembled from a header prefix plus one or more borrowed payload fragments,
// in a single vectored write (writev with a leading virtio_net_hdr).
// This lets the coalescer avoid copying payload bytes between the caller's decrypt buffer and the TUN.
// Backends without GSO support do not implement this interface and coalescing is skipped.
//
// hdr contains the IPv4/IPv6 header prefix (mutable: callers will have filled in total length and IP csum).
// transportHdr is the TCP or UDP header
// (mutable: the L4 checksum field must hold the pseudo-header partial, single-fold not inverted, per virtio NEEDS_CSUM semantics).
// pays are non-overlapping payload fragments whose concatenation is the full superpacket payload.
// They are read-only from the writer's perspective and must remain valid until the call returns.
// Every segment in pays except possibly the last must be exactly the same size.
// proto picks the L4 protocol so the writer knows which gsoType / CsumOffset to set.
//
// Callers should also consult CapsProvider (via SupportsGSO) for the per-protocol negotiated capability:
// USO may not have been negotiated even when TSO was.
type GSOWriter interface {
io.Writer
CapsProvider
WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error
}
// SupportsGSO reports whether w implements GSOWriter and the underlying
// queue advertises the negotiated capability for `want`.
func SupportsGSO(w io.Writer, want GSOProto) (GSOWriter, bool) {
gw, ok := w.(GSOWriter)
if !ok {
return nil, false
}
caps := gw.Capabilities()
switch want {
case GSOProtoTCP:
return gw, caps.TSO
case GSOProtoUDP:
return gw, caps.USO
default:
return gw, false
}
}
-417
View File
@@ -1,417 +0,0 @@
//go:build linux && !android
// +build linux,!android
package tio
import (
"context"
"fmt"
"io"
"log/slog"
"os"
"sync/atomic"
"syscall"
"unsafe"
"golang.org/x/sys/unix"
"github.com/slackhq/nebula/overlay/tio/virtio"
)
const maxSuperpacketLen = 65535
// tunRxBufSize is the per-Read worst-case footprint inside rxBuf: one kernel-supplied packet body, which is at most ~64 KiB.
// Segmentation happens at encrypt time on a per-routine MTU-sized scratch
// (see SegmentSuperpacket), so rxBuf only holds raw kernel-supplied bytes.
// We round up to give margin for the drain headroom check below.
const tunRxBufSize = 64 * 1024
// tunRxBufCap is the total size we allocate for the per-reader rx buffer.
// Each drain iteration consumes up to tunRxBufSize of headroom for the kernel-supplied bytes.
// Sized to eight such iterations so a single poll wake can drain several TSO/USO superpackets under bulk load,
// amortizing the wake and giving the sendmmsg planner longer same-destination runs.
// Hold latency stays bounded because listenIn flushes its send batch incrementally rather than only at end-of-drain.
const tunRxBufCap = tunRxBufSize * 8
// tunDrainCap caps how many packets a single Read will accumulate via the post-wake drain loop.
// Sized to soak up a burst of small ACKs while bounding how much work a single caller holds before handing off.
const tunDrainCap = 64
// gsoMaxIovs caps the iovec budget WriteGSO assembles per call:
// 3 fixed entries (virtio_net_hdr, IP hdr, transport hdr), plus up to gsoMaxIovs-3 payload fragments.
// Sized comfortably above the typical kernel GSO segment cap (Linux UDP_GRO is 64)
// so realistic coalesced bursts never touch the limit.
// iovecs are tiny (16 bytes), so the entire scratch is 4 KiB.
// WriteGSO returns an error rather than reallocating when a caller exceeds this budget.
const gsoMaxIovs = 256
// validVnetHdr is the 10-byte virtio_net_hdr we prepend to every non-GSO TUN write.
// Only flag set is VIRTIO_NET_HDR_F_DATA_VALID. Note the tun write path
// (__virtio_net_hdr_to_skb) ignores this bit — only the virtio-net driver's RX
// helper honors it — so packets land CHECKSUM_NONE and the stack verifies the
// L4 checksum anyway. What matters here is what the header does NOT say:
// no NEEDS_CSUM, so the kernel is never asked to finish a checksum.
// All packets that reach the plain Write paths already carry a valid L4 checksum.
var validVnetHdr = [virtio.Size]byte{unix.VIRTIO_NET_HDR_F_DATA_VALID}
// Offload wraps a TUN file descriptor with poll-based reads. The FD provided will be changed to non-blocking.
// A shared eventfd allows Close to wake all readers blocked in poll.
//
// Field order is deliberate: the read-mostly fds and the writer-owned GSO scratch fill
// the first cache line, and the state the reader mutates per packet (rxOff, pending,
// readIovs) all sits after it, so per-packet reader stores never invalidate the line
// concurrent Write callers load fd from.
type Offload struct {
fd int
shutdownFd int
// usoEnabled records whether the kernel agreed to TUN_F_USO* on this FD,
// so writers can decide whether emitting GSO_UDP_L4 superpackets is safe.
usoEnabled bool
closed atomic.Bool
// gsoHdrBuf is a per-queue 10-byte scratch for the virtio_net_hdr emitted
// by WriteGSO. Kept separate from the read-only package-level validVnetHdr
// so non-GSO Writes can ship that constant directly while WriteGSO
// rewrites this scratch on every call.
gsoHdrBuf [virtio.Size]byte
// gsoIovs is the writev iovec scratch for WriteGSO. Pre-sized to
// gsoMaxIovs at construction; never grown. WriteGSO returns an error
// (and drops the call) if a caller hands it more fragments than fit.
gsoIovs []unix.Iovec
rxBuf []byte // backing store for kernel-handed packets read this drain
rxOff int // cursor into rxBuf for the current Read drain
pending []Packet // packets returned from the most recent Read
// readVnetScratch holds the 10-byte virtio_net_hdr split off the front of
// every TUN read via readv(2). Decoupling the header from the packet body
// lets us read the body directly into rxBuf at the current rxOff with
// no userspace copy on the GSO_NONE fast path.
readVnetScratch [virtio.Size]byte
// readIovs is the readv(2) iovec scratch wired once at construction,
// iovec[0] points at readVnetScratch
// iovec[1].Base/Len is updated per read to address the current rxBuf slot.
readIovs [2]unix.Iovec
// l is only consulted on the rare bad-vnet-header drop path; it lives
// after the hot state on purpose. May be nil (tests); drops go unlogged then.
l *slog.Logger
}
func newOffload(fd int, shutdownFd int, usoEnabled bool, l *slog.Logger) (*Offload, error) {
if err := unix.SetNonblock(fd, true); err != nil {
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
}
out := &Offload{
fd: fd,
shutdownFd: shutdownFd,
usoEnabled: usoEnabled,
closed: atomic.Bool{},
l: l,
rxBuf: make([]byte, tunRxBufCap),
gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs),
}
out.gsoIovs[0].Base = &out.gsoHdrBuf[0]
out.gsoIovs[0].SetLen(virtio.Size)
// readIovs[0] is wired once to the virtio_net_hdr scratch; per-read we
// only repoint readIovs[1] at the next rxBuf slot (see readPacket).
out.readIovs[0].Base = &out.readVnetScratch[0]
out.readIovs[0].SetLen(virtio.Size)
return out, nil
}
func (r *Offload) blockOnRead() error {
return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLIN)
}
func (r *Offload) blockOnWrite() error {
return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLOUT)
}
// readPacket issues a single readv(2), splitting the virtio_net_hdr off into readVnetScratch
// and reading the packet body directly into rxBuf at the current rxOff.
// Returns the body length (zero virtio header bytes, just the IP packet/superpacket).
// block controls whether EAGAIN is retried via poll: the initial read of a drain blocks; subsequent drain reads do not.
func (r *Offload) readPacket(block bool) (int, error) {
for {
r.readIovs[1].Base = &r.rxBuf[r.rxOff]
r.readIovs[1].SetLen(len(r.rxBuf) - r.rxOff)
n, _, errno := syscall.Syscall(unix.SYS_READV, uintptr(r.fd), uintptr(unsafe.Pointer(&r.readIovs[0])), uintptr(len(r.readIovs)))
if errno == 0 {
if int(n) < virtio.Size {
return 0, fmt.Errorf("tun read shorter than virtio_net_hdr: %d bytes", n)
}
return int(n) - virtio.Size, nil
}
if errno == unix.EAGAIN {
if !block {
return 0, errno
}
if err := r.blockOnRead(); err != nil {
return 0, err
}
continue
}
if errno == unix.EINTR {
continue
}
if errno == unix.EBADF {
return 0, os.ErrClosed
}
return 0, errno
}
}
// Read returns one or more packets from the tun.
// Each Packet either carries a single ready-to-use IP datagram (GSO zero) or a TSO/USO superpacket plus the GSOInfo a caller needs to segment it (see SegmentSuperpacket).
// The first read blocks via poll; once the fd is known readable we drain additional packets non-blocking until:
// - the kernel queue is empty (EAGAIN)
// - we've collected tunDrainCap packets,
// - or we're out of rxBuf headroom.
//
// This amortizes the poll wake over bursts of small packets (e.g. TCP ACKs).
// Packet.Bytes slices point into the Offload's internal buffer and are only valid until the next Read or Close on this Queue.
func (r *Offload) Read() ([]Packet, error) {
r.pending = r.pending[:0]
r.rxOff = 0
// Initial (blocking) read.
// Retry on decode errors so a single bad packet does not stall the reader.
for {
n, err := r.readPacket(true)
if err != nil {
return nil, err
}
if err := r.decodeRead(n); err != nil {
// Drop and read again. A bad packet should not kill the reader,
// but a systematic decode failure must not be invisible either.
r.logDroppedRead(err)
continue
}
break
}
// Drain: non-blocking reads until the kernel queue is empty, the drain
// cap is reached, or rxBuf no longer has room for another worst-case
// kernel-supplied packet (tunRxBufSize).
for len(r.pending) < tunDrainCap && tunRxBufCap-r.rxOff >= tunRxBufSize {
n, err := r.readPacket(false)
if err != nil {
// EAGAIN / EINTR / anything else: stop draining. We already
// have a valid batch from the first read.
break
}
if n <= 0 {
break
}
if err := r.decodeRead(n); err != nil {
// Drop this packet and stop the drain; we'd rather hand off
// what we have than keep spinning here.
r.logDroppedRead(err)
break
}
}
return r.pending, nil
}
// logDroppedRead reports a tun packet dropped for a bad/unsupported virtio
// header. Debug-gated so the happy path never pays for attribute assembly.
func (r *Offload) logDroppedRead(err error) {
if r.l != nil && r.l.Enabled(context.Background(), slog.LevelDebug) {
r.l.Debug("dropping tun packet with bad virtio header", "error", err)
}
}
// decodeRead processes the packet sitting in rxBuf at rxOff (length pktLen).
// The bytes stay in rxBuf:
// - for GSO_NONE we slice them as a regular IP datagram (running finishChecksum if NEEDS_CSUM is set);
// - for TSO/USO superpackets we attach the corrected GSO metadata, so the caller can segment lazily at encrypt time.
//
// rxOff advances by pktLen on success
func (r *Offload) decodeRead(pktLen int) error {
if pktLen <= 0 {
return fmt.Errorf("short tun read: %d", pktLen)
}
var hdr virtio.Hdr
hdr.Decode(r.readVnetScratch[:])
body := r.rxBuf[r.rxOff : r.rxOff+pktLen]
if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_NONE {
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
if err := virtio.FinishChecksum(body, hdr); err != nil {
return err
}
}
r.pending = append(r.pending, Packet{Bytes: body})
r.rxOff += pktLen
return nil
}
if err := virtio.CheckValid(body, hdr); err != nil {
return err
}
if err := virtio.CorrectHdrLen(body, &hdr); err != nil {
return err
}
proto, err := protoFromGSOType(hdr.GSOType())
if err != nil {
return err
}
r.pending = append(r.pending, Packet{
Bytes: body,
GSO: GSOInfo{
Size: hdr.GSOSize,
HdrLen: hdr.HdrLen,
CsumStart: hdr.CsumStart,
Proto: proto,
},
})
r.rxOff += pktLen
return nil
}
func (r *Offload) Write(buf []byte) (int, error) {
if len(buf) == 0 {
return 0, nil
}
iovs := [2]unix.Iovec{
{Base: &validVnetHdr[0]},
{Base: &buf[0]},
}
iovs[0].SetLen(virtio.Size)
iovs[1].SetLen(len(buf))
return r.rawWrite(unsafe.Slice(&iovs[0], 2))
}
func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
for {
n, _, errno := syscall.Syscall(unix.SYS_WRITEV, uintptr(r.fd), uintptr(unsafe.Pointer(&iovs[0])), uintptr(len(iovs)))
if errno == 0 {
if int(n) < virtio.Size {
return 0, io.ErrShortWrite
}
return int(n) - virtio.Size, nil
}
if errno == unix.EAGAIN {
if err := r.blockOnWrite(); err != nil {
return 0, err
}
continue
}
if errno == unix.EINTR {
continue
}
if errno == unix.EBADF {
return 0, os.ErrClosed
}
return 0, errno
}
}
// Capabilities reports the offload features negotiated for this Queue. TSO
// is always true for Offload (we only construct it on IFF_VNET_HDR FDs);
// USO is true only when the kernel agreed to TUN_F_USO4|6 at open time (Linux ≥ 6.2).
func (r *Offload) Capabilities() Capabilities {
return Capabilities{TSO: true, USO: r.usoEnabled}
}
func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error {
if len(pays) == 0 {
// There are no payload fragments. There is nothing to send.
return nil
}
var csumOff uint16 // csumOff is the offset of the L4 checksum field in transportHdr
switch proto {
case GSOProtoUDP:
csumOff = 6
case GSOProtoTCP:
csumOff = 16
default:
return fmt.Errorf("unknown GSO proto: %d", proto)
}
// Incorrect geometry must cause an error, not a silent drop.
// No sane packet should ever make it inside this branch.
if len(hdr) == 0 || len(transportHdr) < int(csumOff)+2 {
return fmt.Errorf("tio: WriteGSO header too short: ip=%d transport=%d (csum field at %d)", len(hdr), len(transportHdr), csumOff)
}
// Make the iovec array: [virtio_hdr, hdr, transportHdr, pays...].
// The constructor attaches r.gsoIovs[0] to gsoHdrBuf. That entry does not change.
need := 3 + len(pays)
if need > cap(r.gsoIovs) {
return fmt.Errorf("tio: WriteGSO needs %d iovecs but cap is %d", need, cap(r.gsoIovs))
}
r.gsoIovs = r.gsoIovs[:need]
r.gsoIovs[1].Base = &hdr[0]
r.gsoIovs[1].SetLen(len(hdr))
r.gsoIovs[2].Base = &transportHdr[0]
r.gsoIovs[2].SetLen(len(transportHdr))
segSize := len(pays[0])
total := len(hdr) + len(transportHdr)
for i, p := range pays {
if len(p) == 0 {
// The coalescers route zero-payload packets down the non-GSO path,
// so an empty fragment means the caller's accounting is broken.
return fmt.Errorf("tio: WriteGSO empty payload fragment %d of %d", i, len(pays))
} else if len(p) > segSize || (len(p) < segSize && i != len(pays)-1) {
// all segments must be the same size, except for the last one
return fmt.Errorf("tio: WriteGSO fragment %d is %dB, want %dB segments (only the last may be shorter)", i, len(p), segSize)
}
total += len(p)
r.gsoIovs[3+i].Base = &p[0]
r.gsoIovs[3+i].SetLen(len(p))
}
// This check keeps `total` in the uint16 range. Anything larger would wrap around and cause trouble.
if total > maxSuperpacketLen {
return fmt.Errorf("tio: WriteGSO superpacket %dB exceeds %d", total, maxSuperpacketLen)
}
// A single segment ships as a plain checksummed packet (GSO_NONE, size 0).
// Multiple segments carry the real GSO type and segSize, which the loop
// above verified is the size of every fragment except possibly the last.
gsoType := uint8(unix.VIRTIO_NET_HDR_GSO_NONE)
if len(pays) > 1 {
gsoType = gsoTypeFromProto(proto, hdr[0]>>4)
if gsoType == unix.VIRTIO_NET_HDR_GSO_NONE {
// gsoTypeFromProto only yields GSO_NONE for a bogus IP version nibble.
// A multi-segment superpacket must carry a real GSO type, or the kernel would deliver it as a single jumbo packet.
return fmt.Errorf("tio: WriteGSO IP version %d is not GSO-capable", hdr[0]>>4)
}
}
var gsoSize uint16
if gsoType != unix.VIRTIO_NET_HDR_GSO_NONE {
gsoSize = uint16(segSize)
}
virtio.EncodeHeader(
r.gsoHdrBuf[:],
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
gsoType, /*gsoType*/
uint16(len(hdr)+len(transportHdr)), /*hdrLen*/
gsoSize, /*gsoSize*/
uint16(len(hdr)), /*csumStart*/
csumOff, /*csumOffset*/
)
_, err := r.rawWrite(r.gsoIovs)
return err
}
func (r *Offload) Close() error {
if r.closed.Swap(true) {
return nil
}
// shutdownFd is owned by the container, so we should not close it
// Close the underlying fd but do NOT null r.fd: a reader may still be loading it in readPacket, and mutating the field would race that load.
// That reader gets EBADF -> os.ErrClosed on its next syscall. A reader already parked in
// poll is NOT woken by this close; only the QueueSet's shutdown eventfd wake does that (see Queue.Close docs).
// closed.Swap already guarantees we only close once.
return unix.Close(r.fd)
}
-118
View File
@@ -1,118 +0,0 @@
//go:build linux && !android
// +build linux,!android
package tio
import (
"fmt"
"os"
"sync/atomic"
"golang.org/x/sys/unix"
)
type Poll struct {
fd int
shutdownFd int
closed atomic.Bool
readBuf []byte
batchRet [1]Packet
}
// newPoll wraps an existing tun fd.
// On failure it does NOT close fd: the caller owns fd and is the sole closer
// (see pollQueueSet.Add callers in overlay/tun_linux.go, which unix.Close on Add error).
// This matches the newOffload convention and keeps closes at exactly one on every path.
func newPoll(fd int, shutdownFd int) (*Poll, error) {
if err := unix.SetNonblock(fd, true); err != nil {
return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err)
}
out := &Poll{
fd: fd,
shutdownFd: shutdownFd,
readBuf: make([]byte, 65535), // largest possible size Linux permits
}
return out, nil
}
// blockOnRead waits until the Poll fd is readable or shutdown has been signaled.
// Returns os.ErrClosed if Close was called.
func (t *Poll) blockOnRead() error {
return blockOn(int32(t.fd), int32(t.shutdownFd), unix.POLLIN)
}
func (t *Poll) blockOnWrite() error {
return blockOn(int32(t.fd), int32(t.shutdownFd), unix.POLLOUT)
}
// TODO: port Offload's post-wake drain loop here so one poll wake amortizes
// over a burst (up to tunDrainCap packets) instead of paying a syscall and a
// wake per packet. Hosts on the TUNSETOFFLOAD-failure fallback or a tun.fd
// config currently lose that batching. blockOn and the EAGAIN plumbing are
// already shared; kept one-packet-per-Read for now to preserve behavior.
func (t *Poll) Read() ([]Packet, error) {
n, err := t.readOne(t.readBuf)
if err != nil {
return nil, err
}
t.batchRet[0] = Packet{Bytes: t.readBuf[:n]}
return t.batchRet[:], nil
}
func (t *Poll) readOne(to []byte) (int, error) {
for {
n, errno := unix.Read(t.fd, to)
if errno == nil {
return n, nil
}
switch errno {
case unix.EAGAIN:
if err := t.blockOnRead(); err != nil {
return 0, err
}
case unix.EINTR:
// retry
case unix.EBADF:
return 0, os.ErrClosed
default:
return 0, errno
}
}
}
// Write is safe for concurrent use
func (t *Poll) Write(from []byte) (int, error) {
for {
n, errno := unix.Write(t.fd, from)
if errno == nil {
return n, nil
}
switch errno {
case unix.EAGAIN:
if err := t.blockOnWrite(); err != nil {
return 0, err
}
case unix.EINTR:
// retry
case unix.EBADF:
return 0, os.ErrClosed
default:
return 0, errno
}
}
}
func (t *Poll) Close() error {
if t.closed.Swap(true) {
return nil
}
// shutdownFd is owned by the container, so we should not close it
// Close the underlying fd but do NOT null t.fd: a reader may still be loading it in readOne, and mutating the field would race that load.
// That reader gets EBADF -> os.ErrClosed on its next syscall. A reader already parked in
// poll is NOT woken by this close; only the QueueSet's shutdown eventfd wake does that (see Queue.Close docs).
// closed.Swap already guarantees we only close once.
return unix.Close(t.fd)
}
-228
View File
@@ -1,228 +0,0 @@
//go:build linux && !android && !e2e_testing
// +build linux,!android,!e2e_testing
package tio
import (
"errors"
"log/slog"
"os"
"sync"
"testing"
"time"
"github.com/stretchr/testify/require"
"golang.org/x/sys/unix"
)
// newReadPipe returns a read fd. The matching write fd is registered for cleanup.
// The caller takes ownership of the read fd (pass it into a QueueSet).
func newReadPipe(t *testing.T) int {
t.Helper()
var fds [2]int
if err := unix.Pipe2(fds[:], unix.O_CLOEXEC); err != nil {
t.Fatalf("pipe2: %v", err)
}
t.Cleanup(func() { _ = unix.Close(fds[1]) })
return fds[0]
}
func TestPoll_WakeForShutdown_WakesFriends(t *testing.T) {
pipe1 := newReadPipe(t)
pipe2 := newReadPipe(t)
parent, err := NewPollQueueSet()
require.NoError(t, err)
require.NoError(t, parent.Add(pipe1))
require.NoError(t, parent.Add(pipe2))
t.Cleanup(func() {
_ = unix.Close(pipe1)
_ = unix.Close(pipe2)
})
readers := parent.Queues()
errs := make([]error, len(readers))
var wg sync.WaitGroup
for i, r := range readers {
wg.Add(1)
go func(i int, r Queue) {
defer wg.Done()
_, errs[i] = r.Read()
}(i, r)
}
time.Sleep(50 * time.Millisecond)
if err := parent.Close(); err != nil {
t.Fatalf("Close: %v", err)
}
done := make(chan struct{})
go func() { wg.Wait(); close(done) }()
select {
case <-done:
case <-time.After(2 * time.Second):
t.Fatal("readers did not wake")
}
for i, err := range errs {
if !errors.Is(err, os.ErrClosed) {
t.Errorf("reader %d: expected os.ErrClosed, got %v", i, err)
}
}
}
// TestPoll_ConcurrentWrite_NoRace hammers a single Poll queue from two writer
// goroutines while a reader drains the other end of the pipe. The writers
// overflow the pipe buffer, so both repeatedly park in blockOnWrite at the same
// time — the exact scenario that raced on the old shared writePoll member
// array. Run under -race; a shared-array regression trips the detector here.
func TestPoll_ConcurrentWrite_NoRace(t *testing.T) {
var fds [2]int
require.NoError(t, unix.Pipe2(fds[:], unix.O_CLOEXEC))
readFd, writeFd := fds[0], fds[1]
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
require.NoError(t, err)
t.Cleanup(func() { _ = unix.Close(shutdownFd) })
p, err := newPoll(writeFd, shutdownFd)
require.NoError(t, err)
const writers = 2
const perWriter = 4000
payload := make([]byte, 100)
total := writers * perWriter * len(payload)
// Reader: drain the read end (blocking) until every writer's bytes are
// consumed, so the writers keep making progress rather than wedging on a
// permanently full pipe.
readDone := make(chan struct{})
go func() {
defer close(readDone)
buf := make([]byte, 4096)
got := 0
for got < total {
n, rerr := unix.Read(readFd, buf)
got += n
if rerr != nil {
if rerr == unix.EINTR {
continue
}
return
}
if n == 0 { // EOF
return
}
}
}()
var wg sync.WaitGroup
for w := 0; w < writers; w++ {
wg.Add(1)
go func() {
defer wg.Done()
for i := 0; i < perWriter; i++ {
if _, werr := p.Write(payload); werr != nil {
t.Errorf("write: %v", werr)
return
}
}
}()
}
wg.Wait()
select {
case <-readDone:
case <-time.After(10 * time.Second):
t.Fatal("reader did not drain")
}
require.NoError(t, p.Close())
_ = unix.Close(readFd)
}
// TestPoll_NewPoll_DoesNotCloseFdOnFailure pins the ownership rule: when
// newPoll fails, it must leave fd open so the caller (pollQueueSet.Add's
// callers in tun_linux.go) is the sole closer. If newPoll also closed fd,
// the poll path would double-close on Add error. We force the failure with
// an O_PATH descriptor: fcntl(F_SETFL) — which SetNonblock performs — is not
// permitted on O_PATH fds and fails with EBADF, while the fd itself stays
// open so we can observe that newPoll left it alone.
func TestPoll_NewPoll_DoesNotCloseFdOnFailure(t *testing.T) {
fd, err := unix.Open("/", unix.O_PATH|unix.O_CLOEXEC, 0)
require.NoError(t, err)
t.Cleanup(func() { _ = unix.Close(fd) })
p, err := newPoll(fd, 1)
require.Error(t, err, "SetNonblock on an O_PATH fd should fail")
require.Nil(t, p)
// If newPoll had closed fd, F_GETFD would report it closed. It staying
// open proves newPoll left the fd for the caller to close exactly once.
require.True(t, fdOpen(t, fd), "newPoll must not close fd on failure; caller is the sole closer")
}
func TestPoll_Close_Idempotent(t *testing.T) {
tf, err := newPoll(newReadPipe(t), 1)
require.NoError(t, err)
if err := tf.Close(); err != nil {
t.Fatalf("first Close: %v", err)
}
if err := tf.Close(); err != nil {
t.Fatalf("second Close should be a no-op, got %v", err)
}
}
// fdOpen reports whether fd currently refers to an open file description.
// A closed (or never-allocated) fd makes F_GETFD fail with EBADF.
func fdOpen(t *testing.T, fd int) bool {
t.Helper()
_, err := unix.FcntlInt(uintptr(fd), unix.F_GETFD, 0)
if err == nil {
return true
}
if errors.Is(err, unix.EBADF) {
return false
}
t.Fatalf("unexpected fcntl(F_GETFD) error on fd %d: %v", fd, err)
return false
}
// TestPollQueueSet_Close_ClosesShutdownFd is the regression test for the
// leaked shutdown eventfd: the container that owns shutdownFd must close it in
// Close, and a second Close must be a safe no-op.
func TestPollQueueSet_Close_ClosesShutdownFd(t *testing.T) {
qs, err := NewPollQueueSet()
require.NoError(t, err)
c, ok := qs.(*pollQueueSet)
require.True(t, ok)
require.NoError(t, qs.Add(newReadPipe(t)))
shutdownFd := c.shutdownFd
require.True(t, fdOpen(t, shutdownFd), "shutdown eventfd should be open before Close")
require.NoError(t, qs.Close())
require.False(t, fdOpen(t, shutdownFd), "shutdown eventfd should be closed after Close")
// Second Close must not touch fds (shutdownFd is now -1) and must return nil.
require.NoError(t, qs.Close())
}
// TestOffloadQueueSet_Close_ClosesShutdownFd mirrors the poll regression test
// for the GSO/offload queueset.
func TestOffloadQueueSet_Close_ClosesShutdownFd(t *testing.T) {
qs, err := NewOffloadQueueSet(false, slog.New(slog.DiscardHandler))
require.NoError(t, err)
c, ok := qs.(*offloadQueueSet)
require.True(t, ok)
require.NoError(t, qs.Add(newReadPipe(t)))
shutdownFd := c.shutdownFd
require.True(t, fdOpen(t, shutdownFd), "shutdown eventfd should be open before Close")
require.NoError(t, qs.Close())
require.False(t, fdOpen(t, shutdownFd), "shutdown eventfd should be closed after Close")
// Second Close must not touch fds (shutdownFd is now -1) and must return nil.
require.NoError(t, qs.Close())
}
-60
View File
@@ -1,60 +0,0 @@
//go:build linux && !android
// +build linux,!android
package tio
import (
"fmt"
"golang.org/x/sys/unix"
"github.com/slackhq/nebula/overlay/tio/virtio"
)
// protoFromGSOType maps a virtio_net_hdr gsoType to the GSOProto value the
// segment-time helpers use. Returns an error for GSO_NONE or any unknown
// value. The caller should only invoke this on a confirmed superpacket.
func protoFromGSOType(t uint8) (GSOProto, error) {
switch t &^ unix.VIRTIO_NET_HDR_GSO_ECN {
case unix.VIRTIO_NET_HDR_GSO_TCPV4, unix.VIRTIO_NET_HDR_GSO_TCPV6:
return GSOProtoTCP, nil
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
return GSOProtoUDP, nil
default:
return 0, fmt.Errorf("unsupported virtio gso type: %d", t)
}
}
// gsoTypeFromProto is the reverse of protoFromGSOType
func gsoTypeFromProto(proto GSOProto, ipVer uint8) uint8 {
switch {
case proto == GSOProtoUDP && (ipVer == 4 || ipVer == 6):
return unix.VIRTIO_NET_HDR_GSO_UDP_L4
case ipVer == 6:
return unix.VIRTIO_NET_HDR_GSO_TCPV6
case ipVer == 4:
return unix.VIRTIO_NET_HDR_GSO_TCPV4
default:
return unix.VIRTIO_NET_HDR_GSO_NONE
}
}
// SegmentSuperpacket invokes fn once per segment of pkt.
// For non-GSO pkts fn is called once with pkt.Bytes.
// For GSO/USO superpackets, fn is called once per segment with a slice of pkt.Bytes holding that segment's plaintext
// (a freshly-patched L3+L4 header sliced in front of the original payload chunk).
// This slicing is destructive: pkt is consumed by this call.
// Aborts and returns the first error from fn or from per-segment construction.
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
if !pkt.GSO.IsSuperpacket() {
return fn(pkt.Bytes)
}
switch pkt.GSO.Proto {
case GSOProtoTCP:
return virtio.SegmentTCP(pkt.Bytes, pkt.GSO.HdrLen, pkt.GSO.CsumStart, pkt.GSO.Size, fn)
case GSOProtoUDP:
return virtio.SegmentUDP(pkt.Bytes, pkt.GSO.HdrLen, pkt.GSO.CsumStart, pkt.GSO.Size, fn)
default:
return fmt.Errorf("unsupported gso proto: %d", pkt.GSO.Proto)
}
}
File diff suppressed because it is too large Load Diff
-75
View File
@@ -1,75 +0,0 @@
//go:build linux && !android
// +build linux,!android
package virtio
import (
"encoding/binary"
"golang.org/x/sys/unix"
)
// Size is the on-wire length of struct virtio_net_hdr the kernel
// prepends/expects on a TUN opened with IFF_VNET_HDR (TUNSETVNETHDRSZ
// not set).
const Size = 10
// Hdr is the Go view of the legacy virtio_net_hdr.
type Hdr struct {
Flags uint8
gsoType uint8 //private to avoid mistakes wrt the 0x80 VIRTIO_NET_HDR_GSO_ECN flag, ORed with the other "GSO types"
HdrLen uint16
GSOSize uint16
CsumStart uint16
CsumOffset uint16
}
func NewHeader(flags, gsoType uint8, hdrLen, gsoSize, csumStart, csumOffset uint16) Hdr {
return Hdr{
Flags: flags,
gsoType: gsoType,
HdrLen: hdrLen,
GSOSize: gsoSize,
CsumStart: csumStart,
CsumOffset: csumOffset,
}
}
// Decode reads a virtio_net_hdr in host byte order (TUN default; we never
// call TUNSETVNETLE so the kernel matches our endianness).
func (h *Hdr) Decode(b []byte) {
h.Flags = b[0]
h.gsoType = b[1]
h.HdrLen = binary.NativeEndian.Uint16(b[2:4])
h.GSOSize = binary.NativeEndian.Uint16(b[4:6])
h.CsumStart = binary.NativeEndian.Uint16(b[6:8])
h.CsumOffset = binary.NativeEndian.Uint16(b[8:10])
}
func EncodeHeader(b []byte, flags, gsoType uint8, hdrLen, gsoSize, csumStart, csumOffset uint16) {
b[0] = flags
b[1] = gsoType
binary.NativeEndian.PutUint16(b[2:4], hdrLen)
binary.NativeEndian.PutUint16(b[4:6], gsoSize)
binary.NativeEndian.PutUint16(b[6:8], csumStart)
binary.NativeEndian.PutUint16(b[8:10], csumOffset)
}
// Encode is the inverse of Decode: writes the virtio_net_hdr fields into b
// (must be at least Size bytes). Used to emit a TSO superpacket on egress.
func (h *Hdr) Encode(b []byte) {
EncodeHeader(b, h.Flags, h.gsoType, h.HdrLen, h.GSOSize, h.CsumStart, h.CsumOffset)
}
// GSOType returns gsoType with the ECN-flag masked out
func (h *Hdr) GSOType() uint8 {
return h.gsoType &^ unix.VIRTIO_NET_HDR_GSO_ECN
}
func (h *Hdr) HasECNFlag() bool {
return h.gsoType&unix.VIRTIO_NET_HDR_GSO_ECN != 0
}
func (h *Hdr) SetGSOType(x uint8) {
h.gsoType = x
}
-3
View File
@@ -1,3 +0,0 @@
//go:build !linux || android
package virtio
-441
View File
@@ -1,441 +0,0 @@
//go:build linux && !android
// +build linux,!android
// Package virtio implements the pure validation, header-correction, and
// per-segment slicing logic for kernel-supplied TSO/USO superpackets on
// IFF_VNET_HDR TUN devices. It is FD-free and depends only on the byte
// layout of the virtio_net_hdr and the IP/TCP/UDP headers it describes,
// so it can be unit-tested in isolation from the tio Queue runtime.
package virtio
import (
"encoding/binary"
"errors"
"fmt"
"golang.org/x/sys/unix"
"github.com/slackhq/nebula/overlay/checksum"
)
// Protocol header size bounds used to validate / cap kernel-supplied offsets.
const (
ipv4HeaderMinLen = 20 // IHL=5, no options
ipv4HeaderMaxLen = 60 // IHL=15, max options
ipv6FixedLen = 40 // IPv6 base header; extensions would extend this
tcpHeaderMinLen = 20 // data-offset=5, no options
tcpHeaderMaxLen = 60 // data-offset=15, max options
)
// maxSegHdrLen bounds the L3+L4 header we snapshot before stamping each segment.
// The largest header the segmenter supports is IPv4 (max IHL 60) plus TCP (max data-offset 60) = 120 bytes
const maxSegHdrLen = ipv4HeaderMaxLen + tcpHeaderMaxLen // 120
// Byte offsets inside an IPv4 header.
const (
ipv4TotalLenOff = 2
ipv4IDOff = 4
ipv4ChecksumOff = 10
ipv4SrcOff = 12
ipv4AddrsEnd = 20 // end of dst address (ipv4SrcOff + 2*4)
)
// Byte offsets inside an IPv6 header.
const (
ipv6PayloadLenOff = 4
ipv6SrcOff = 8
ipv6AddrsEnd = 40 // end of dst address (ipv6SrcOff + 2*16)
)
// Byte offsets inside a TCP header (relative to its start, i.e. csumStart).
const (
tcpSeqOff = 4
tcpDataOffOff = 12 // upper nibble is header len in 32-bit words
tcpFlagsOff = 13
tcpChecksumOff = 16
)
// UDP header is fixed at 8 bytes: {sport, dport, length, checksum}.
const (
udpHeaderLen = 8
udpLengthOff = 4
udpChecksumOff = 6
)
var errPacketTooShort = errors.New("packet too short")
// tcpFinPshMask is cleared on every segment except the last of a TSO burst.
const tcpFinPshMask = 0x09 // FIN(0x01) | PSH(0x08)
// tcpCwrFlag is cleared on every segment except the first.
// Per RFC 3168 §6.1.2 the CWR bit signals a one-shot transition (the sender just halved its window)
// and must appear on the first segment of a TSO burst only.
const tcpCwrFlag = 0x80
// CheckValid rejects packets whose virtio_net_hdr/IP combination would
// cause a downstream miscompute. The TUN should never emit RSC_INFO and
// the GSO type must agree with the IP version nibble.
func CheckValid(pkt []byte, hdr Hdr) error {
if hdr.Flags&unix.VIRTIO_NET_HDR_F_RSC_INFO != 0 {
return fmt.Errorf("virtio RSC_INFO flag not supported on TUN reads")
}
if len(pkt) < ipv4HeaderMinLen {
return errPacketTooShort
}
ipVersion := pkt[0] >> 4
if ipVersion == 6 && len(pkt) < ipv6FixedLen {
return errPacketTooShort
}
gsoType := hdr.GSOType()
if gsoType != unix.VIRTIO_NET_HDR_GSO_NONE && hdr.GSOSize == 0 {
// A GSO type with no segment size would dodge IsSuperpacket() downstream and
// travel as a plain jumbo datagram with an unfinished checksum.
return fmt.Errorf("virtio GSO type %#x with zero gso_size", hdr.gsoType)
}
if hdr.HasECNFlag() && !(gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV4 || gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV6) {
return fmt.Errorf("virtio GSO_ECN qualifier on non-TCP GSO type %#x", hdr.gsoType)
}
switch gsoType {
case unix.VIRTIO_NET_HDR_GSO_TCPV4:
if ipVersion != 4 {
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
}
case unix.VIRTIO_NET_HDR_GSO_TCPV6:
if ipVersion != 6 {
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
}
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
// USO carries either v4 or v6; the leading nibble disambiguates.
if !(ipVersion == 4 || ipVersion == 6) {
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
}
default:
if !(ipVersion == 6 || ipVersion == 4) {
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
}
}
return nil
}
// CorrectHdrLen rewrites hdr.HdrLen based on the actual transport header length read out of pkt.
// The kernel's hdr.HdrLen on the FORWARD path can be the length of the entire first packet, so we don't trust it.
func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
// Thank you wireguard-go for documenting these edge-cases
// Don't trust hdr.hdrLen from the kernel as it can be equal to the length
// of the entire first packet when the kernel is handling it as part of a FORWARD path.
// Instead, parse the transport header length and add it onto csumStart, which is synonymous for IP header length.
if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
hdr.HdrLen = hdr.CsumStart + 8
} else {
if len(pkt) <= int(hdr.CsumStart+tcpDataOffOff) {
return errors.New("packet is too short")
}
tcpHLen := uint16(pkt[hdr.CsumStart+tcpDataOffOff] >> 4 * 4)
if tcpHLen < tcpHeaderMinLen || tcpHLen > tcpHeaderMaxLen {
return fmt.Errorf("tcp header len is invalid: %d", tcpHLen)
}
hdr.HdrLen = hdr.CsumStart + tcpHLen
}
if len(pkt) < int(hdr.HdrLen) {
return fmt.Errorf("length of packet (%d) < virtioNetHdr.HdrLen (%d)", len(pkt), hdr.HdrLen)
}
if hdr.HdrLen < hdr.CsumStart {
return fmt.Errorf("virtioNetHdr.HdrLen (%d) < virtioNetHdr.CsumStart (%d)", hdr.HdrLen, hdr.CsumStart)
}
cSumAt := int(hdr.CsumStart + hdr.CsumOffset)
if cSumAt+1 >= len(pkt) {
return fmt.Errorf("end of checksum offset (%d) exceeds packet length (%d)", cSumAt+1, len(pkt))
}
return nil
}
// segCount returns how many segments a payload of payLen bytes splits into at gsoSize,
// with a floor of one so a header-only superpacket still yields a single segment.
func segCount(payLen, gsoSize int) int {
n := (payLen + gsoSize - 1) / gsoSize
if n == 0 {
return 1
}
return n
}
// basePseudoSum folds the part of the L4 pseudo-header sum that is identical
// for every segment: the source and destination addresses plus the protocol
// number. The per-segment L4 length is added by the caller inside the loop.
func basePseudoSum(pkt []byte, isV4 bool, proto uint32) uint32 {
if isV4 {
return uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0)) + proto
}
return uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0)) + proto
}
// baseIPv4HdrSum folds the IPv4 header checksum over the fields that stay constant across segments.
// csumStart is the L3 header length, which bounds a valid IHL.
func baseIPv4HdrSum(pkt []byte, csumStart int) (uint32, error) {
ihl := int(pkt[0]&0x0f) * 4
if ihl < ipv4HeaderMinLen || ihl > csumStart {
return 0, fmt.Errorf("bad IPv4 IHL: %d", ihl)
}
// total_len, the ID, and the checksum field itself are excluded: all three are rewritten per segment.
sum := uint32(checksum.Checksum(pkt[:ihl], 0))
sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4TotalLenOff : ipv4TotalLenOff+2]))
sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4ChecksumOff : ipv4ChecksumOff+2]))
sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2]))
sum = (sum & 0xffff) + (sum >> 16)
sum = (sum & 0xffff) + (sum >> 16)
return sum, nil
}
// baseTCPHdrSum folds the TCP header checksum over everything the segment loop does not rewrite
func baseTCPHdrSum(pkt []byte, csumStart, headerLen int) uint32 {
seq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
flags := uint16(pkt[csumStart+tcpFlagsOff])
sum := uint32(checksum.Checksum(pkt[csumStart:headerLen], 0))
sum += uint32(^uint16(seq >> 16))
sum += uint32(^uint16(seq))
sum += uint32(^flags)
sum += uint32(^binary.BigEndian.Uint16(pkt[csumStart+tcpChecksumOff : csumStart+tcpChecksumOff+2]))
sum = (sum & 0xffff) + (sum >> 16)
sum = (sum & 0xffff) + (sum >> 16)
return sum
}
// SegmentTCP walks a TSO superpacket pkt, yielding each segment as a slice into pkt.
// Per-segment plaintext is laid out by stamping a copy of the original L3+L4 header into pkt at offset i*gsoSize,
// where it sits immediately before that segment's payload chunk in the original buffer.
// pkt is consumed by this call and must not be inspected by the caller after the final yield.
func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
if gsoSizeU == 0 {
return fmt.Errorf("gso_size is zero")
}
if csumStartU == 0 {
return fmt.Errorf("csum_start is zero")
}
headerLen := int(hdrLenU)
csumStart := int(csumStartU)
if headerLen > maxSegHdrLen {
return fmt.Errorf("header len %d exceeds max %d", headerLen, maxSegHdrLen)
}
isV4 := pkt[0]>>4 == 4
tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4
payLen := len(pkt) - headerLen
gsoSize := int(gsoSizeU)
numSeg := segCount(payLen, gsoSize)
origSeq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
origFlags := pkt[csumStart+tcpFlagsOff]
baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_TCP)
baseTcpHdrSum := baseTCPHdrSum(pkt, csumStart, headerLen)
var origIPID uint16
var baseIPHdrSum uint32
if isV4 {
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
var err error
// TSO bumps the ID per segment, so it stays out of the base sum.
baseIPHdrSum, err = baseIPv4HdrSum(pkt, csumStart)
if err != nil {
return err
}
}
// Snapshot the pristine L3+L4 header once. '
// Every segment's header is stamped from this copy, so overlapping stamps (gsoSize < headerLen) can never corrupt the source.
var savedHdr [maxSegHdrLen]byte
copy(savedHdr[:headerLen], pkt[:headerLen])
for i := 0; i < numSeg; i++ {
segStart := i * gsoSize
segEnd := segStart + gsoSize
if segEnd > payLen {
segEnd = payLen
}
segPayLen := segEnd - segStart
segLen := headerLen + segPayLen
headerOff := i * gsoSize
// Stamp the header into place immediately before this segment's payload, sourced from the snapshot.
// The per-segment patches below overwrite the variable fields. (seq/flags/cksum/totalLen/id)
if i > 0 {
// Iter 0's header is already at pkt[:headerLen] (identical to savedHdr), so only i >= 1 needs the stamp
copy(pkt[headerOff:headerOff+headerLen], savedHdr[:headerLen])
}
seg := pkt[headerOff : headerOff+segLen]
segSeq := origSeq + uint32(segStart)
segFlags := origFlags
if i != 0 {
segFlags &^= tcpCwrFlag
}
if i != numSeg-1 {
segFlags &^= tcpFinPshMask
}
totalLen := segLen
if isV4 {
segID := origIPID + uint16(i)
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID)
ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID)
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
} else {
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
}
binary.BigEndian.PutUint32(seg[csumStart+tcpSeqOff:csumStart+tcpSeqOff+4], segSeq)
seg[csumStart+tcpFlagsOff] = segFlags
tcpLen := tcpHdrLen + segPayLen
// Payload bytes still live at their original offset in pkt.
// The header slide above only writes into pkt[i*GSOSize : i*GSOSize+header], which is the tail of seg_{i-1}'s payload (already consumed)
// and never overlaps seg_i's own payload at pkt[header+i*GSOSize : header+(i+1)*GSOSize].
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
wide := uint64(baseTcpHdrSum) + uint64(paySum) + uint64(baseProtoSum)
wide += uint64(segSeq) + uint64(segFlags) + uint64(tcpLen)
wide = (wide & 0xffffffff) + (wide >> 32)
wide = (wide & 0xffffffff) + (wide >> 32)
binary.BigEndian.PutUint16(seg[csumStart+tcpChecksumOff:csumStart+tcpChecksumOff+2], foldComplement(uint32(wide)))
if err := yield(seg); err != nil {
return err
}
}
return nil
}
// SegmentUDP walks a USO superpacket, stamping a per-segment-patched copy of the original L3+L4 header
// into pkt at offset i*GSOSize and yielding pkt[i*GSOSize:i*GSOSize+segLen] to the caller.
// Per-segment patches are total_len + IPv4 csum (or IPv6 payload_len) plus the UDP length and checksum.
// pkt is consumed destructively.
func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
if gsoSizeU == 0 {
return fmt.Errorf("gso_size is zero")
}
if csumStartU == 0 {
return fmt.Errorf("csum_start is zero")
}
isV4 := pkt[0]>>4 == 4
headerLen := int(hdrLenU)
csumStart := int(csumStartU)
if headerLen > maxSegHdrLen {
return fmt.Errorf("header len %d exceeds max %d", headerLen, maxSegHdrLen)
}
if headerLen-csumStart != udpHeaderLen {
return fmt.Errorf("udp header len mismatch: %d", headerLen-csumStart)
}
payLen := len(pkt) - headerLen
gsoSize := int(gsoSizeU)
numSeg := segCount(payLen, gsoSize)
baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_UDP)
var origIPID uint16
var baseIPHdrSum uint32
if isV4 {
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
var err error
// Software UDP GSO bumps the ID per segment just like TSO
// (inet_gso_segment's fixed-ID case is TCP-only), so it stays out of the base sum.
baseIPHdrSum, err = baseIPv4HdrSum(pkt, csumStart)
if err != nil {
return err
}
}
// Snapshot the pristine L3+L4 header once and stamp every segment from it
var savedHdr [maxSegHdrLen]byte
copy(savedHdr[:headerLen], pkt[:headerLen])
for i := 0; i < numSeg; i++ {
segStart := i * gsoSize
segEnd := segStart + gsoSize
if segEnd > payLen {
segEnd = payLen
}
segPayLen := segEnd - segStart
segLen := headerLen + segPayLen
headerOff := i * gsoSize
if i > 0 {
copy(pkt[headerOff:headerOff+headerLen], savedHdr[:headerLen])
}
seg := pkt[headerOff : headerOff+segLen]
totalLen := segLen
udpLen := udpHeaderLen + segPayLen
if isV4 {
segID := origIPID + uint16(i)
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID)
ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID)
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
} else {
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
}
binary.BigEndian.PutUint16(seg[csumStart+udpLengthOff:csumStart+udpLengthOff+2], uint16(udpLen))
// Sum the UDP header (length just written, checksum zeroed) together with
// this segment's payload in one pass, seeded with the pseudo-header sum.
seg[csumStart+udpChecksumOff], seg[csumStart+udpChecksumOff+1] = 0, 0
pseudo := baseProtoSum + uint32(udpLen)
pseudo = (pseudo & 0xffff) + (pseudo >> 16)
pseudo = (pseudo & 0xffff) + (pseudo >> 16)
csum := ^checksum.Checksum(seg[csumStart:], uint16(pseudo))
if csum == 0 {
csum = 0xffff
}
binary.BigEndian.PutUint16(seg[csumStart+udpChecksumOff:csumStart+udpChecksumOff+2], csum)
if err := yield(seg); err != nil {
return err
}
}
return nil
}
// FinishChecksum computes the L4 checksum for a non-GSO packet that the kernel handed us with NEEDS_CSUM set.
// CsumStart / CsumOffset point at the 16-bit checksum field.
// We zero it, fold a full sum from the partial one that the kernel provided, and store the result.
func FinishChecksum(seg []byte, hdr Hdr) error {
cs := int(hdr.CsumStart)
co := int(hdr.CsumOffset)
if cs+co+2 > len(seg) {
return fmt.Errorf("csum offsets out of range: start=%d offset=%d len=%d", cs, co, len(seg))
}
// The kernel stores a partial pseudo-header sum at [cs+co:]; sum over the
// L4 region starting at cs, folding the prior partial in as the seed.
partial := binary.BigEndian.Uint16(seg[cs+co : cs+co+2])
seg[cs+co] = 0
seg[cs+co+1] = 0
csum := ^checksum.Checksum(seg[cs:], partial)
// RFC 768: UDP transmits a computed zero as all ones, since all-zero is the reserved "no checksum" value.
if co == udpChecksumOff && csum == 0 {
csum = 0xffff
}
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], csum)
return nil
}
// foldComplement folds a 32-bit one's-complement partial sum to 16 bits and
// complements it, yielding the on-wire Internet checksum value.
func foldComplement(sum uint32) uint16 {
sum = (sum & 0xffff) + (sum >> 16)
sum = (sum & 0xffff) + (sum >> 16)
return ^uint16(sum)
}
-602
View File
@@ -1,602 +0,0 @@
//go:build linux && !android
// +build linux,!android
package virtio
import (
"bytes"
"encoding/binary"
"testing"
"golang.org/x/sys/unix"
"github.com/slackhq/nebula/overlay/checksum"
)
// verifyChecksum confirms that the one's-complement sum across b, seeded with
// a folded pseudo-header sum, equals all-ones (a valid on-wire checksum).
// A corrupted header stamped into a segment makes this fail even when the
// checksum field itself was computed from the (pristine) base sums, because
// the bytes the receiver would sum no longer match what was checksummed.
func verifyChecksum(b []byte, pseudo uint16) bool {
return checksum.Checksum(b, pseudo) == 0xffff
}
// pseudoHeaderIPv4 folds the TCP/UDP pseudo-header sum from a segment's own
// address and length fields, used to independently verify its L4 checksum.
func pseudoHeaderIPv4(src, dst []byte, proto byte, l4Len int) uint16 {
s := uint32(checksum.Checksum(src, 0)) + uint32(checksum.Checksum(dst, 0))
s += uint32(proto) + uint32(l4Len)
s = (s & 0xffff) + (s >> 16)
s = (s & 0xffff) + (s >> 16)
return uint16(s)
}
// buildTCPv4Super constructs a synthetic IPv4/TCP TSO superpacket with a
// payload of payLen bytes and returns it alongside the header fields the
// segmenter needs. The header is a fixed 40 bytes (20 IPv4 + 20 TCP).
func buildTCPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
const ipLen = 20
const tcpLen = 20
pkt = make([]byte, ipLen+tcpLen+payLen)
// IPv4 header.
pkt[0] = 0x45 // version 4, IHL 5
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+tcpLen+payLen))
binary.BigEndian.PutUint16(pkt[4:6], 0x4242) // ID
pkt[8] = 64 // TTL
pkt[9] = unix.IPPROTO_TCP
copy(pkt[12:16], []byte{10, 0, 0, 1}) // src
copy(pkt[16:20], []byte{10, 0, 0, 2}) // dst
// TCP header.
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
binary.BigEndian.PutUint16(pkt[22:24], 80) // dport
binary.BigEndian.PutUint32(pkt[24:28], 10000) // seq
binary.BigEndian.PutUint32(pkt[28:32], 20000) // ack
pkt[32] = 0x50 // data offset 5 words
pkt[33] = 0x18 // ACK | PSH
binary.BigEndian.PutUint16(pkt[34:36], 65535) // window
for i := 0; i < payLen; i++ {
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
}
return pkt, ipLen + tcpLen, ipLen
}
// buildUDPv4Super constructs a synthetic IPv4/UDP USO superpacket with a
// payload of payLen bytes. Header is a fixed 28 bytes (20 IPv4 + 8 UDP).
func buildUDPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
const ipLen = 20
const udpLen = 8
pkt = make([]byte, ipLen+udpLen+payLen)
pkt[0] = 0x45
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+udpLen+payLen))
binary.BigEndian.PutUint16(pkt[4:6], 0x4242)
pkt[8] = 64
pkt[9] = unix.IPPROTO_UDP
copy(pkt[12:16], []byte{10, 0, 0, 1})
copy(pkt[16:20], []byte{10, 0, 0, 2})
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
for i := 0; i < payLen; i++ {
pkt[ipLen+udpLen+i] = byte(i & 0xff)
}
return pkt, ipLen + udpLen, ipLen
}
// collectTCP segments a fresh copy of pkt and returns each segment as an
// independent slice so assertions can run after segmentation completes.
func collectTCP(t *testing.T, pkt []byte, hdrLen, csumStart, gsoSize uint16) [][]byte {
t.Helper()
work := append([]byte(nil), pkt...)
var out [][]byte
err := SegmentTCP(work, hdrLen, csumStart, gsoSize, func(seg []byte) error {
out = append(out, append([]byte(nil), seg...))
return nil
})
if err != nil {
t.Fatalf("SegmentTCP: %v", err)
}
return out
}
func collectUDP(t *testing.T, pkt []byte, hdrLen, csumStart, gsoSize uint16) [][]byte {
t.Helper()
work := append([]byte(nil), pkt...)
var out [][]byte
err := SegmentUDP(work, hdrLen, csumStart, gsoSize, func(seg []byte) error {
out = append(out, append([]byte(nil), seg...))
return nil
})
if err != nil {
t.Fatalf("SegmentUDP: %v", err)
}
return out
}
// TestSegmentTCPHeaderNotCorrupted is the regression test for the in-place
// header-slide bug: when gsoSize < headerLen the old code stamped each
// segment's header from pkt[:headerLen], which had already been overwritten
// by the previous segment's overlapping stamp, so segments 2..n carried a
// corrupted header (garbage src/dst/ports/seq). Every segment must instead
// carry the ORIGINAL constant header fields with correct per-segment seq.
func TestSegmentTCPHeaderNotCorrupted(t *testing.T) {
const origSeq = 10000
cases := []struct {
name string
payLen int
gsoSize uint16
}{
// gsoSize (8) < headerLen (40): the bug's trigger. Even split.
{"small-gso-even", 40, 8},
// gsoSize (8) < headerLen (40) with a short final segment.
{"small-gso-odd-tail", 44, 8},
// gsoSize (100) >= headerLen (40): the normal path, must still work.
{"normal-gso", 250, 100},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
pkt, hdrLen, csumStart := buildTCPv4Super(tc.payLen)
gso := int(tc.gsoSize)
wantSeg := (tc.payLen + gso - 1) / gso
segs := collectTCP(t, pkt, hdrLen, csumStart, tc.gsoSize)
if len(segs) != wantSeg {
t.Fatalf("got %d segments, want %d", len(segs), wantSeg)
}
off := 0
for i, seg := range segs {
// Constant header fields must be identical to the original in
// EVERY segment. These are exactly the bytes the old code
// corrupted in segments 2..n.
if got := seg[0]; got != 0x45 {
t.Errorf("seg %d: version/IHL byte=%#x want 0x45", i, got)
}
if seg[9] != unix.IPPROTO_TCP {
t.Errorf("seg %d: proto=%d want %d", i, seg[9], unix.IPPROTO_TCP)
}
if !bytes.Equal(seg[12:16], []byte{10, 0, 0, 1}) {
t.Errorf("seg %d: src=%v want [10 0 0 1]", i, seg[12:16])
}
if !bytes.Equal(seg[16:20], []byte{10, 0, 0, 2}) {
t.Errorf("seg %d: dst=%v want [10 0 0 2]", i, seg[16:20])
}
if sport := binary.BigEndian.Uint16(seg[20:22]); sport != 12345 {
t.Errorf("seg %d: sport=%d want 12345", i, sport)
}
if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 80 {
t.Errorf("seg %d: dport=%d want 80", i, dport)
}
if ack := binary.BigEndian.Uint32(seg[28:32]); ack != 20000 {
t.Errorf("seg %d: ack=%d want 20000", i, ack)
}
if seg[32] != 0x50 {
t.Errorf("seg %d: data-offset byte=%#x want 0x50", i, seg[32])
}
// Per-segment seq must advance by the payload offset.
segStart := i * gso
if seq := binary.BigEndian.Uint32(seg[24:28]); seq != uint32(origSeq+segStart) {
t.Errorf("seg %d: seq=%d want %d", i, seq, origSeq+segStart)
}
// Payload bytes must be the original contiguous slice.
segPayLen := len(seg) - int(hdrLen)
wantPay := make([]byte, segPayLen)
for k := 0; k < segPayLen; k++ {
wantPay[k] = byte((off + k) & 0xff)
}
if !bytes.Equal(seg[hdrLen:], wantPay) {
t.Errorf("seg %d: payload mismatch", i)
}
off += segPayLen
// End-to-end: the stamped header must checksum-verify. A
// corrupted header fails here because the written checksum was
// derived from the pristine header.
if !verifyChecksum(seg[:20], 0) {
t.Errorf("seg %d: bad IPv4 header checksum", i)
}
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_TCP, len(seg)-20)
if !verifyChecksum(seg[20:], psum) {
t.Errorf("seg %d: bad TCP checksum", i)
}
}
})
}
}
// TestCorrectHdrLenChecksumBound guards the checksum-field bounds check in
// CorrectHdrLen. The checksum field sits at CsumStart+CsumOffset, so the check
// must be computed from CsumStart+CsumOffset — NOT CsumStart+CsumStart, a
// regression that doubled CsumStart and thus over-tightened the bound (since
// CsumOffset, 6 for UDP / 16 for TCP, is always < CsumStart >= 20). That bogus
// bound spuriously rejected valid small USO superpackets in decodeRead.
func TestCorrectHdrLenChecksumBound(t *testing.T) {
// A valid IPv4 USO superpacket: 20B IPv4 + 8B UDP + two 6-byte segments
// (payload 12) = 40 bytes total. CsumStart=20, CsumOffset=6, so the UDP
// checksum field lives at bytes 26..27, comfortably inside the 40-byte
// packet. The OLD formula computed cSumAt = CsumStart+CsumStart = 40 and
// rejected on cSumAt+1 (41) >= len(pkt) (40); the fix (CsumStart+CsumOffset
// = 26) accepts. This case FAILS against the CsumStart+CsumStart regression.
t.Run("valid-small-uso-accepted", func(t *testing.T) {
pkt, _, csumStart := buildUDPv4Super(12) // total len 40
hdr := NewHeader(
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
0, /*hdrLen*/
6, /*gsoSize: two 6-byte segments*/
csumStart, /*csumStart*/
6, /*csumOffset*/
)
if err := CorrectHdrLen(pkt, &hdr); err != nil {
t.Fatalf("CorrectHdrLen rejected a valid 40-byte USO superpacket: %v", err)
}
if hdr.HdrLen != csumStart+udpHeaderLen {
t.Errorf("HdrLen = %d, want %d", hdr.HdrLen, csumStart+udpHeaderLen)
}
})
// A genuinely-too-short packet: CsumStart=20, CsumOffset=6 means the
// checksum field would end at byte 27, but the packet is only 25 bytes
// (CsumStart+CsumOffset+2 = 28 > 25). CorrectHdrLen must still reject it.
t.Run("too-short-rejected", func(t *testing.T) {
pkt := make([]byte, 25)
pkt[0] = 0x45 // IPv4, IHL 5
hdr := NewHeader(
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
0, /*hdrLen*/
6, /*gsoSize*/
20, /*csumStart*/
6, /*csumOffset*/
)
if err := CorrectHdrLen(pkt, &hdr); err == nil {
t.Fatalf("CorrectHdrLen accepted a too-short (25-byte) packet")
}
})
}
// TestSegmentUDPHeaderNotCorrupted is the USO counterpart: SegmentUDP performs
// the same header stamp and must be correct when gsoSize < headerLen.
func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
cases := []struct {
name string
payLen int
gsoSize uint16
}{
{"small-gso-even", 40, 8},
{"small-gso-odd-tail", 44, 8},
{"normal-gso", 250, 100},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
pkt, hdrLen, csumStart := buildUDPv4Super(tc.payLen)
gso := int(tc.gsoSize)
wantSeg := (tc.payLen + gso - 1) / gso
segs := collectUDP(t, pkt, hdrLen, csumStart, tc.gsoSize)
if len(segs) != wantSeg {
t.Fatalf("got %d segments, want %d", len(segs), wantSeg)
}
off := 0
for i, seg := range segs {
if got := seg[0]; got != 0x45 {
t.Errorf("seg %d: version/IHL byte=%#x want 0x45", i, got)
}
if seg[9] != unix.IPPROTO_UDP {
t.Errorf("seg %d: proto=%d want %d", i, seg[9], unix.IPPROTO_UDP)
}
if !bytes.Equal(seg[12:16], []byte{10, 0, 0, 1}) {
t.Errorf("seg %d: src=%v want [10 0 0 1]", i, seg[12:16])
}
if !bytes.Equal(seg[16:20], []byte{10, 0, 0, 2}) {
t.Errorf("seg %d: dst=%v want [10 0 0 2]", i, seg[16:20])
}
if sport := binary.BigEndian.Uint16(seg[20:22]); sport != 12345 {
t.Errorf("seg %d: sport=%d want 12345", i, sport)
}
if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 53 {
t.Errorf("seg %d: dport=%d want 53", i, dport)
}
// Software UDP GSO bumps the IPv4 ID per segment just like TSO
// (inet_gso_segment's fixed-ID case is TCP-only).
if id := binary.BigEndian.Uint16(seg[4:6]); id != 0x4242+uint16(i) {
t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242+uint16(i))
}
segPayLen := len(seg) - int(hdrLen)
if udpLen := binary.BigEndian.Uint16(seg[24:26]); udpLen != uint16(8+segPayLen) {
t.Errorf("seg %d: udp len=%d want %d", i, udpLen, 8+segPayLen)
}
wantPay := make([]byte, segPayLen)
for k := 0; k < segPayLen; k++ {
wantPay[k] = byte((off + k) & 0xff)
}
if !bytes.Equal(seg[hdrLen:], wantPay) {
t.Errorf("seg %d: payload mismatch", i)
}
off += segPayLen
if !verifyChecksum(seg[:20], 0) {
t.Errorf("seg %d: bad IPv4 header checksum", i)
}
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_UDP, len(seg)-20)
if !verifyChecksum(seg[20:], psum) {
t.Errorf("seg %d: bad UDP checksum", i)
}
}
})
}
}
// buildUDPv4Single constructs a single IPv4/UDP datagram with the checksum field pre-loaded with the folded
// pseudo-header sum, which is how a NEEDS_CSUM (CHECKSUM_PARTIAL) packet arrives from the tun.
func buildUDPv4Single(payload []byte) (pkt []byte, hdr Hdr) {
const ipLen, udpLen = 20, 8
pkt = make([]byte, ipLen+udpLen+len(payload))
pkt[0] = 0x45
binary.BigEndian.PutUint16(pkt[2:4], uint16(len(pkt)))
pkt[8] = 64
pkt[9] = unix.IPPROTO_UDP
copy(pkt[12:16], []byte{10, 0, 0, 1})
copy(pkt[16:20], []byte{10, 0, 0, 2})
binary.BigEndian.PutUint16(pkt[ipLen:ipLen+2], 12345)
binary.BigEndian.PutUint16(pkt[ipLen+2:ipLen+4], 53)
binary.BigEndian.PutUint16(pkt[ipLen+4:ipLen+6], uint16(udpLen+len(payload)))
copy(pkt[ipLen+udpLen:], payload)
pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, udpLen+len(payload))
binary.BigEndian.PutUint16(pkt[ipLen+udpChecksumOff:ipLen+udpChecksumOff+2], pseudo)
return pkt, NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, ipLen, udpChecksumOff)
}
// TestFinishChecksumUDPZeroStoresAllOnes pins RFC 768: a UDP checksum that computes to zero goes on the wire as
// 0xffff, because all-zero is the reserved "no checksum" encoding that IPv6 rejects outright.
func TestFinishChecksumUDPZeroStoresAllOnes(t *testing.T) {
var payload []byte
for i := 0; i < 0x10000; i++ {
p := []byte{byte(i >> 8), byte(i)}
pkt, hdr := buildUDPv4Single(p)
cs, co := int(hdr.CsumStart), int(hdr.CsumOffset)
partial := binary.BigEndian.Uint16(pkt[cs+co : cs+co+2])
pkt[cs+co], pkt[cs+co+1] = 0, 0
if ^checksum.Checksum(pkt[cs:], partial) == 0 {
payload = p
break
}
}
if payload == nil {
t.Fatal("no 2-byte payload produced a zero checksum")
}
pkt, hdr := buildUDPv4Single(payload)
if err := FinishChecksum(pkt, hdr); err != nil {
t.Fatalf("FinishChecksum: %v", err)
}
off := int(hdr.CsumStart) + int(hdr.CsumOffset)
if got := binary.BigEndian.Uint16(pkt[off : off+2]); got != 0xffff {
t.Fatalf("stored %#04x, want 0xffff for a UDP checksum that folds to zero", got)
}
}
// TestFinishChecksumTCPZeroPreserved is the other half of the protocol split: zero is a legal TCP checksum and must
// be stored as-is, so the UDP rewrite must not fire on a TCP csum_offset.
func TestFinishChecksumTCPZeroPreserved(t *testing.T) {
const cs, co = 20, tcpChecksumOff
// Seed the partial so the completed checksum lands on zero, the case the UDP path rewrites.
seg := make([]byte, cs+co+2)
for i := range seg[cs:] {
seg[cs+i] = byte(i * 7)
}
var partial uint16
for i := 0; i <= 0xffff; i++ {
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], uint16(i))
probe := append([]byte(nil), seg...)
probe[cs+co], probe[cs+co+1] = 0, 0
if ^checksum.Checksum(probe[cs:], uint16(i)) == 0 {
partial = uint16(i)
break
}
}
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], partial)
hdr := NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, cs, co)
if err := FinishChecksum(seg, hdr); err != nil {
t.Fatalf("FinishChecksum: %v", err)
}
if got := binary.BigEndian.Uint16(seg[cs+co : cs+co+2]); got != 0 {
t.Fatalf("stored %#04x, want 0x0000 (zero is a legal TCP checksum)", got)
}
}
// TestFinishChecksumUDPValidates confirms the ordinary path still produces a checksum a receiver accepts.
func TestFinishChecksumUDPValidates(t *testing.T) {
payload := []byte("the definitive tun offloads branch")
pkt, hdr := buildUDPv4Single(payload)
if err := FinishChecksum(pkt, hdr); err != nil {
t.Fatalf("FinishChecksum: %v", err)
}
pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, 8+len(payload))
if !verifyChecksum(pkt[hdr.CsumStart:], pseudo) {
t.Fatal("completed UDP checksum does not validate")
}
}
// TestCheckValidMasksGSOECN: GSO_ECN is a qualifier bit the kernel ORs
// into gso_type for TSO superpackets with CWR set. CheckValid must
// validate an ECN-qualified type as its base type — previously TCPV4|ECN
// fell into the default case and skipped the IP-version agreement check.
// The qualifier is TCP-only, so it must be rejected on UDP_L4.
func TestCheckValidMasksGSOECN(t *testing.T) {
v4pkt, _, _ := buildTCPv4Super(100)
v6pkt := make([]byte, len(v4pkt))
copy(v6pkt, v4pkt)
v6pkt[0] = 0x60 // claim IPv6
cases := []struct {
name string
pkt []byte
gsoType uint8
wantErr bool
}{
{"tcpv4-ecn-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, false},
{"tcpv4-ecn-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
{"tcpv6-ecn-v4-mismatch", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV6 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
{"udp-l4-ecn-rejected", v4pkt, unix.VIRTIO_NET_HDR_GSO_UDP_L4 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
{"tcpv4-plain-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, false},
{"tcpv4-plain-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, true},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
err := CheckValid(tc.pkt, NewHeader(0, tc.gsoType, 0, 100, 0, 0))
if tc.wantErr && err == nil {
t.Errorf("CheckValid(gsoType=%#x) = nil, want error", tc.gsoType)
}
if !tc.wantErr && err != nil {
t.Errorf("CheckValid(gsoType=%#x) = %v, want nil", tc.gsoType, err)
}
})
}
}
// TestCheckValidRejectsZeroGSOSize: a GSO-typed header with gso_size=0 must be
// rejected. It would produce a Packet whose GSOInfo.IsSuperpacket() is false,
// dodging both segmentation and FinishChecksum on its way downstream.
func TestCheckValidRejectsZeroGSOSize(t *testing.T) {
v4pkt, _, _ := buildTCPv4Super(100)
if err := CheckValid(v4pkt, NewHeader(0, unix.VIRTIO_NET_HDR_GSO_TCPV4, 0, 0, 0, 0)); err == nil {
t.Fatal("CheckValid accepted a GSO-typed header with gso_size=0")
}
}
// TestFoldComplementMatchesReference checks the segmenter's fold-and-invert
// against an independent RFC 1071 reference fold, hitting the carry edge
// cases (values whose first fold produces another carry).
func TestFoldComplementMatchesReference(t *testing.T) {
refFold := func(s uint64) uint16 {
for s>>16 != 0 {
s = s&0xffff + s>>16
}
return uint16(s)
}
cases := []uint32{
0, 1, 0xffff,
0x10000, // single carry
0x1fffe, // 0xffff + 0xffff: carry produces another 0xffff
0xffff0000, // high half only
0xfffeffff, // first fold yields another carry
0xffffffff, // worst case
}
for _, c := range cases {
if got, want := foldComplement(c), ^refFold(uint64(c)); got != want {
t.Errorf("foldComplement(%#x) = %#x, want %#x", c, got, want)
}
}
}
// referenceBaseIPv4HdrSum and referenceBaseTCPHdrSum are the straightforward
// implementations that baseIPv4HdrSum/baseTCPHdrSum replaced: copy the header
// into scratch, zero the fields the segment loop rewrites, sum. The production
// versions instead sum in place and subtract those fields via one's-complement
// arithmetic, which is faster but far less obvious — particularly for the TCP
// flags byte, which is only half of a 16-bit word. These references exist so
// that trade is checked rather than asserted.
func referenceBaseIPv4HdrSum(pkt []byte, ihl int) uint32 {
var ipTmp [ipv4HeaderMaxLen]byte
copy(ipTmp[:ihl], pkt[:ihl])
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0
return uint32(checksum.Checksum(ipTmp[:ihl], 0))
}
func referenceBaseTCPHdrSum(pkt []byte, csumStart, headerLen int) uint32 {
tcpLen := headerLen - csumStart
var tmp [tcpHeaderMaxLen]byte
copy(tmp[:tcpLen], pkt[csumStart:headerLen])
tmp[tcpSeqOff], tmp[tcpSeqOff+1], tmp[tcpSeqOff+2], tmp[tcpSeqOff+3] = 0, 0, 0, 0
tmp[tcpFlagsOff] = 0
tmp[tcpChecksumOff], tmp[tcpChecksumOff+1] = 0, 0
return uint32(checksum.Checksum(tmp[:tcpLen], 0))
}
// randSeed is a tiny deterministic PRNG so this test needs no imports beyond
// what the file already has and reproduces identically on every run.
func randByte(state *uint32) byte {
*state = *state*1664525 + 1013904223
return byte(*state >> 24)
}
func TestBaseSumsMatchZeroingReference(t *testing.T) {
state := uint32(12345)
t.Run("ipv4", func(t *testing.T) {
for ihl := ipv4HeaderMinLen; ihl <= ipv4HeaderMaxLen; ihl += 4 {
for iter := 0; iter < 5000; iter++ {
pkt := make([]byte, ihl)
for i := range pkt {
pkt[i] = randByte(&state)
}
pkt[0] = byte(0x40 | (ihl / 4))
want := referenceBaseIPv4HdrSum(pkt, ihl)
got, err := baseIPv4HdrSum(pkt, ihl)
if err != nil {
t.Fatalf("ihl=%d: %v", ihl, err)
}
// Compare the value that reaches the wire: the raw partial
// sums may legally differ by one's-complement -0 vs +0.
for _, tl := range []uint32{20, 1500, 65535} {
for _, id := range []uint32{0, 0x4242, 0xffff} {
if a, b := foldComplement(want+tl+id), foldComplement(got+tl+id); a != b {
t.Fatalf("ihl=%d tl=%d id=%d: %#04x != %#04x", ihl, tl, id, a, b)
}
}
}
}
}
})
t.Run("tcp", func(t *testing.T) {
const csumStart = 20
for dataOff := 5; dataOff <= 15; dataOff++ {
tcpLen := dataOff * 4
headerLen := csumStart + tcpLen
for iter := 0; iter < 5000; iter++ {
pkt := make([]byte, headerLen+64)
for i := range pkt {
pkt[i] = randByte(&state)
}
pkt[0] = 0x45
pkt[csumStart+tcpDataOffOff] = byte(dataOff << 4)
want := referenceBaseTCPHdrSum(pkt, csumStart, headerLen)
got := baseTCPHdrSum(pkt, csumStart, headerLen)
for _, seq := range []uint32{0, 1, 0x4242_4242, 0xffff_ffff} {
for _, fl := range []uint32{0x00, 0x10, 0x18, 0x19, 0xff} {
for _, l4 := range []uint32{20, 1460, 65535} {
a := foldComplement(want + seq + fl + l4)
b := foldComplement(got + seq + fl + l4)
if a != b {
t.Fatalf("dataOff=%d seq=%#x fl=%#x l4=%d: %#04x != %#04x",
dataOff, seq, fl, l4, a, b)
}
}
}
}
}
}
})
}

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