mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 23:56:57 +02:00
Compare commits
37 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 7fe4fab167 | |||
| 37b924945d | |||
| 5631346b07 | |||
| 97eb3c635a | |||
| 05f7923860 | |||
| 49028cb755 | |||
| adf71d1458 | |||
| cefda6524c | |||
| d0f14de739 | |||
| 9e7646ee62 | |||
| 6e6cfc89db | |||
| 3719f135e3 | |||
| 2724b4a96c | |||
| e386e290ab | |||
| 0a44376403 | |||
| 9e61269935 | |||
| a0836aa819 | |||
| cf2800f7bd | |||
| fc89b9d14c | |||
| 3e5326e60a | |||
| 840841a53c | |||
| cfb4ab24c3 | |||
| 67f9cfad91 | |||
| 1c78f5e500 | |||
| ad8ff2e45e | |||
| c13a1ff4ec | |||
| 8b66ebaf72 | |||
| bbdaf5c3f3 | |||
| a515aeaf39 | |||
| 0f6f14eaf6 | |||
| dafdd34af9 | |||
| f6791130df | |||
| 145a6267fa | |||
| 7716f1da23 | |||
| beb1d7d89f | |||
| fa0d593f28 | |||
| ff48040c78 |
@@ -161,6 +161,10 @@ bin-pkcs11: BUILD_ARGS += -tags pkcs11
|
|||||||
bin-pkcs11: CGO_ENABLED = 1
|
bin-pkcs11: CGO_ENABLED = 1
|
||||||
bin-pkcs11: bin
|
bin-pkcs11: bin
|
||||||
|
|
||||||
|
# Build with the pprof debug server (serves on :6060). See startPprofServer.
|
||||||
|
debug: BUILD_ARGS += -tags debug
|
||||||
|
debug: bin
|
||||||
|
|
||||||
bin:
|
bin:
|
||||||
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
|
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
|
||||||
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
|
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
|
||||||
@@ -280,5 +284,5 @@ smoke-vagrant/%: bin-docker build/%/nebula
|
|||||||
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
|
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
|
||||||
|
|
||||||
.FORCE:
|
.FORCE:
|
||||||
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
|
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin debug build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
|
||||||
.DEFAULT_GOAL := bin
|
.DEFAULT_GOAL := bin
|
||||||
|
|||||||
+1
-1
@@ -10,7 +10,7 @@ import (
|
|||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
)
|
)
|
||||||
|
|
||||||
const ReplayWindow = 1024
|
const ReplayWindow = 8192
|
||||||
|
|
||||||
type ConnectionState struct {
|
type ConnectionState struct {
|
||||||
eKey noiseutil.CipherState
|
eKey noiseutil.CipherState
|
||||||
|
|||||||
+27
-14
@@ -11,6 +11,7 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
@@ -77,6 +78,7 @@ func newReadyControl(t *testing.T) (*Control, *fakeDevice, *fakeConn) {
|
|||||||
inside: dev,
|
inside: dev,
|
||||||
outside: conn,
|
outside: conn,
|
||||||
writers: []udp.Conn{conn},
|
writers: []udp.Conn{conn},
|
||||||
|
batchers: make([]batch.RxBatcher, 1),
|
||||||
routines: 1,
|
routines: 1,
|
||||||
hostMap: newHostMap(l),
|
hostMap: newHostMap(l),
|
||||||
lightHouse: lh,
|
lightHouse: lh,
|
||||||
@@ -107,7 +109,8 @@ func TestControl_StopBeforeStart(t *testing.T) {
|
|||||||
require.NoError(t, c.Wait())
|
require.NoError(t, c.Wait())
|
||||||
|
|
||||||
// A stopped control can never be started
|
// A stopped control can never be started
|
||||||
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
err := c.Start()
|
||||||
|
require.ErrorIs(t, err, ErrAlreadyStopped)
|
||||||
|
|
||||||
// A second Stop is a harmless no-op
|
// A second Stop is a harmless no-op
|
||||||
c.Stop()
|
c.Stop()
|
||||||
@@ -141,13 +144,16 @@ type fakeConn struct {
|
|||||||
rebinds int
|
rebinds int
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *fakeConn) Rebind() error { c.rebinds++; return nil }
|
func (c *fakeConn) Rebind() error { c.rebinds++; return nil }
|
||||||
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
||||||
func (c *fakeConn) ListenOut(_ udp.EncReader) error { return nil }
|
func (c *fakeConn) ListenOut(_ udp.EncReader, _ func()) error { return nil }
|
||||||
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
|
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
|
||||||
func (c *fakeConn) ReloadConfig(_ *config.C) {}
|
func (c *fakeConn) WriteBatch(_ [][]byte, _ []netip.AddrPort, _ []byte) error {
|
||||||
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
|
return nil
|
||||||
func (c *fakeConn) Close() error { c.closed = true; return nil }
|
}
|
||||||
|
func (c *fakeConn) ReloadConfig(_ *config.C) {}
|
||||||
|
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
|
||||||
|
func (c *fakeConn) Close() error { c.closed = true; return nil }
|
||||||
|
|
||||||
type multiqueueDevice struct {
|
type multiqueueDevice struct {
|
||||||
*fakeDevice
|
*fakeDevice
|
||||||
@@ -171,6 +177,7 @@ func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
|||||||
inside: dev,
|
inside: dev,
|
||||||
outside: conn,
|
outside: conn,
|
||||||
writers: []udp.Conn{conn},
|
writers: []udp.Conn{conn},
|
||||||
|
batchers: make([]batch.RxBatcher, 2),
|
||||||
routines: 2,
|
routines: 2,
|
||||||
l: test.NewLogger(),
|
l: test.NewLogger(),
|
||||||
}
|
}
|
||||||
@@ -185,7 +192,8 @@ func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// The second reader fails to open, everything must be released
|
// The second reader fails to open, everything must be released
|
||||||
require.Error(t, c.Start())
|
err := c.Start()
|
||||||
|
require.Error(t, err)
|
||||||
assert.Equal(t, StateStopped, c.State())
|
assert.Equal(t, StateStopped, c.State())
|
||||||
assert.True(t, dev.closed, "the tun device should have been closed")
|
assert.True(t, dev.closed, "the tun device should have been closed")
|
||||||
assert.True(t, conn.closed, "the udp socket should have been closed")
|
assert.True(t, conn.closed, "the udp socket should have been closed")
|
||||||
@@ -255,15 +263,18 @@ func TestControl_ConcurrentStopAndStart(t *testing.T) {
|
|||||||
// panic and Wait must observe the final state
|
// panic and Wait must observe the final state
|
||||||
require.NoError(t, c.Wait())
|
require.NoError(t, c.Wait())
|
||||||
assert.Equal(t, StateStopped, c.State())
|
assert.Equal(t, StateStopped, c.State())
|
||||||
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
err := c.Start()
|
||||||
|
require.ErrorIs(t, err, ErrAlreadyStopped)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestControl_StartStopLifecycle(t *testing.T) {
|
func TestControl_StartStopLifecycle(t *testing.T) {
|
||||||
c, dev, conn := newReadyControl(t)
|
c, dev, conn := newReadyControl(t)
|
||||||
|
|
||||||
require.NoError(t, c.Start())
|
err := c.Start()
|
||||||
|
require.NoError(t, err)
|
||||||
assert.Equal(t, StateStarted, c.State())
|
assert.Equal(t, StateStarted, c.State())
|
||||||
require.ErrorIs(t, c.Start(), ErrAlreadyStarted)
|
err = c.Start()
|
||||||
|
require.ErrorIs(t, err, ErrAlreadyStarted)
|
||||||
|
|
||||||
// Stop must unpark the reader blocked in the device and release everything
|
// Stop must unpark the reader blocked in the device and release everything
|
||||||
c.Stop()
|
c.Stop()
|
||||||
@@ -274,7 +285,8 @@ func TestControl_StartStopLifecycle(t *testing.T) {
|
|||||||
|
|
||||||
// The reader drained off a closed device, that is not a fatal error
|
// The reader drained off a closed device, that is not a fatal error
|
||||||
require.NoError(t, c.Wait())
|
require.NoError(t, c.Wait())
|
||||||
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
err = c.Start()
|
||||||
|
require.ErrorIs(t, err, ErrAlreadyStopped)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestControl_RebindIsGatedByState(t *testing.T) {
|
func TestControl_RebindIsGatedByState(t *testing.T) {
|
||||||
@@ -284,7 +296,8 @@ func TestControl_RebindIsGatedByState(t *testing.T) {
|
|||||||
c.RebindUDPServer()
|
c.RebindUDPServer()
|
||||||
assert.Equal(t, 0, conn.rebinds, "rebind before start must be a no-op")
|
assert.Equal(t, 0, conn.rebinds, "rebind before start must be a no-op")
|
||||||
|
|
||||||
require.NoError(t, c.Start())
|
err := c.Start()
|
||||||
|
require.NoError(t, err)
|
||||||
c.RebindUDPServer()
|
c.RebindUDPServer()
|
||||||
assert.Equal(t, 1, conn.rebinds, "rebind while started must reach the conn")
|
assert.Equal(t, 1, conn.rebinds, "rebind while started must reach the conn")
|
||||||
|
|
||||||
|
|||||||
+2
-4
@@ -4,15 +4,13 @@
|
|||||||
package e2e
|
package e2e
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"io"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"log/slog"
|
|
||||||
|
|
||||||
"dario.cat/mergo"
|
"dario.cat/mergo"
|
||||||
"github.com/google/gopacket"
|
"github.com/google/gopacket"
|
||||||
"github.com/google/gopacket/layers"
|
"github.com/google/gopacket/layers"
|
||||||
@@ -382,7 +380,7 @@ func getAddrs(ns []netip.Prefix) []netip.Addr {
|
|||||||
func NewTestLogger() *slog.Logger {
|
func NewTestLogger() *slog.Logger {
|
||||||
v := os.Getenv("TEST_LOGS")
|
v := os.Getenv("TEST_LOGS")
|
||||||
if v == "" {
|
if v == "" {
|
||||||
return slog.New(slog.NewTextHandler(io.Discard, nil))
|
return slog.New(slog.DiscardHandler)
|
||||||
}
|
}
|
||||||
|
|
||||||
level := slog.LevelInfo
|
level := slog.LevelInfo
|
||||||
|
|||||||
@@ -0,0 +1,188 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"log/slog"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"golang.org/x/net/ipv4"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestInnerECN(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
pkt []byte
|
||||||
|
want byte
|
||||||
|
}{
|
||||||
|
{"empty", nil, 0},
|
||||||
|
{"v4_NotECT", v4WithToS(0x00), 0x00},
|
||||||
|
{"v4_ECT0", v4WithToS(0x02), 0x02},
|
||||||
|
{"v4_ECT1", v4WithToS(0x01), 0x01},
|
||||||
|
{"v4_CE", v4WithToS(0x03), 0x03},
|
||||||
|
{"v4_DSCP_then_NotECT", v4WithToS(0x88 | 0x00), 0x00},
|
||||||
|
{"v4_DSCP_then_CE", v4WithToS(0x88 | 0x03), 0x03},
|
||||||
|
{"v6_NotECT", v6WithTC(0x00), 0x00},
|
||||||
|
{"v6_ECT0", v6WithTC(0x02), 0x02},
|
||||||
|
{"v6_CE", v6WithTC(0x03), 0x03},
|
||||||
|
{"v6_DSCP_then_CE", v6WithTC(0x88 | 0x03), 0x03},
|
||||||
|
{"unknown_version", []byte{0xa5, 0xff}, 0},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
t.Run(c.name, func(t *testing.T) {
|
||||||
|
got := innerECN(c.pkt)
|
||||||
|
if got != c.want {
|
||||||
|
t.Errorf("innerECN=0x%02x want 0x%02x", got, c.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// v4WithToS returns a 2-byte slice tall enough for innerECN: byte 0 carries
|
||||||
|
// version=4 in the high nibble, byte 1 is the full ToS so we exercise both
|
||||||
|
// the DSCP and ECN portions through the byte 1 mask.
|
||||||
|
func v4WithToS(tos byte) []byte {
|
||||||
|
return []byte{0x45, tos}
|
||||||
|
}
|
||||||
|
|
||||||
|
// v6WithTC builds a 2-byte slice that places a known traffic class value
|
||||||
|
// across bytes 0 (high nibble of TC) and 1 (low nibble of TC). innerECN
|
||||||
|
// extracts ECN as (b[1]>>4)&0x03, which corresponds to TC[1:0].
|
||||||
|
func v6WithTC(tc byte) []byte {
|
||||||
|
return []byte{0x60 | (tc>>4)&0x0f, (tc & 0x0f) << 4}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestApplyOuterECN(t *testing.T) {
|
||||||
|
silent := slog.New(slog.DiscardHandler)
|
||||||
|
hi := &HostInfo{}
|
||||||
|
|
||||||
|
// Build a v4 packet helper with a given inner ECN field.
|
||||||
|
v4 := func(innerECN byte) []byte {
|
||||||
|
// 20-byte minimal IPv4 header with ToS = innerECN (DSCP zeroed).
|
||||||
|
return []byte{
|
||||||
|
0x45, innerECN, 0, 28,
|
||||||
|
0, 0, 0x40, 0,
|
||||||
|
64, 6, 0, 0,
|
||||||
|
10, 0, 0, 1,
|
||||||
|
10, 0, 0, 2,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Build a v6 packet helper with a given inner ECN field. ECN occupies
|
||||||
|
// TC[1:0] which sit at byte 1 mask 0x30.
|
||||||
|
v6 := func(innerECN byte) []byte {
|
||||||
|
// 40-byte minimal IPv6 header with TC[1:0] = innerECN.
|
||||||
|
pkt := make([]byte, 40)
|
||||||
|
pkt[0] = 0x60 // version=6, TC[7:4]=0
|
||||||
|
pkt[1] = (innerECN & 0x03) << 4 // TC[3:0]: low 2 bits = ECN, top 2 = DSCP-low (0)
|
||||||
|
return pkt
|
||||||
|
}
|
||||||
|
|
||||||
|
type cell struct {
|
||||||
|
outer byte
|
||||||
|
inner byte
|
||||||
|
wantECN byte
|
||||||
|
wantSame bool // expect inner unchanged (true => verify the byte didn't move)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RFC 6040 normal-mode combine table. Only outer==CE causes mutation.
|
||||||
|
table := []cell{
|
||||||
|
{ecnNotECT, ecnNotECT, ecnNotECT, true},
|
||||||
|
{ecnNotECT, ecnECT0, ecnECT0, true},
|
||||||
|
{ecnNotECT, ecnECT1, ecnECT1, true},
|
||||||
|
{ecnNotECT, ecnCE, ecnCE, true},
|
||||||
|
|
||||||
|
{ecnECT0, ecnNotECT, ecnNotECT, true},
|
||||||
|
{ecnECT0, ecnECT0, ecnECT0, true},
|
||||||
|
{ecnECT0, ecnECT1, ecnECT1, true},
|
||||||
|
{ecnECT0, ecnCE, ecnCE, true},
|
||||||
|
|
||||||
|
{ecnECT1, ecnNotECT, ecnNotECT, true},
|
||||||
|
{ecnECT1, ecnECT0, ecnECT0, true},
|
||||||
|
{ecnECT1, ecnECT1, ecnECT1, true},
|
||||||
|
{ecnECT1, ecnCE, ecnCE, true},
|
||||||
|
|
||||||
|
{ecnCE, ecnNotECT, ecnNotECT, true}, // legacy: log, leave alone
|
||||||
|
{ecnCE, ecnECT0, ecnCE, false}, // CE folded in
|
||||||
|
{ecnCE, ecnECT1, ecnCE, false},
|
||||||
|
{ecnCE, ecnCE, ecnCE, true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, c := range table {
|
||||||
|
t.Run("v4", func(t *testing.T) {
|
||||||
|
pkt := v4(c.inner)
|
||||||
|
applyOuterECN(pkt, c.outer, hi, silent)
|
||||||
|
got := pkt[1] & 0x03
|
||||||
|
if got != c.wantECN {
|
||||||
|
t.Errorf("v4 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
t.Run("v6", func(t *testing.T) {
|
||||||
|
pkt := v6(c.inner)
|
||||||
|
applyOuterECN(pkt, c.outer, hi, silent)
|
||||||
|
got := (pkt[1] >> 4) & 0x03
|
||||||
|
if got != c.wantECN {
|
||||||
|
t.Errorf("v6 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestApplyOuterECN_IPv4ChecksumStaysValid guards against H1: folding an outer
|
||||||
|
// CE mark into the inner IPv4 ToS byte must keep the IPv4 header checksum valid.
|
||||||
|
// The passthrough emit paths write the packet verbatim, so a stale checksum
|
||||||
|
// turns an underlay congestion mark into packet loss at the receiver.
|
||||||
|
func TestApplyOuterECN_IPv4ChecksumStaysValid(t *testing.T) {
|
||||||
|
silent := slog.New(slog.DiscardHandler)
|
||||||
|
hi := &HostInfo{}
|
||||||
|
|
||||||
|
// 20-byte IPv4 header with DSCP=0x88 and inner ECN = ECT(0). Folding CE
|
||||||
|
// flips only the low two bits of the ToS byte while leaving DSCP intact.
|
||||||
|
pkt := []byte{
|
||||||
|
0x45, 0x88 | ecnECT0, 0, 40,
|
||||||
|
0x1c, 0x46, 0x40, 0x00,
|
||||||
|
64, 6, 0, 0,
|
||||||
|
10, 0, 0, 1,
|
||||||
|
10, 0, 0, 2,
|
||||||
|
}
|
||||||
|
// Stamp a correct header checksum before the fold.
|
||||||
|
binary.BigEndian.PutUint16(pkt[10:12], ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]))
|
||||||
|
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
|
||||||
|
t.Fatal("test setup: initial header checksum invalid")
|
||||||
|
}
|
||||||
|
|
||||||
|
applyOuterECN(pkt, ecnCE, hi, silent)
|
||||||
|
|
||||||
|
// CE folded in, DSCP preserved.
|
||||||
|
if got, want := pkt[1], byte(0x88|ecnCE); got != want {
|
||||||
|
t.Fatalf("ToS after fold = 0x%02x, want 0x%02x", got, want)
|
||||||
|
}
|
||||||
|
// The incremental RFC 1624 update must leave the checksum valid and equal
|
||||||
|
// to a full recompute over the mutated header.
|
||||||
|
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
|
||||||
|
t.Fatalf("IPv4 header checksum invalid after CE fold: 0x%04x", binary.BigEndian.Uint16(pkt[10:12]))
|
||||||
|
}
|
||||||
|
if got, want := binary.BigEndian.Uint16(pkt[10:12]), ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]); got != want {
|
||||||
|
t.Fatalf("checksum = 0x%04x, full recompute = 0x%04x", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ipv4HeaderChecksum computes the RFC 1071 IPv4 header checksum over hdr,
|
||||||
|
// treating the checksum field (bytes 10:12) as zero.
|
||||||
|
func ipv4HeaderChecksum(hdr []byte) uint16 {
|
||||||
|
var sum uint32
|
||||||
|
for i := 0; i+1 < len(hdr); i += 2 {
|
||||||
|
if i == 10 {
|
||||||
|
continue // checksum field
|
||||||
|
}
|
||||||
|
sum += uint32(hdr[i])<<8 | uint32(hdr[i+1])
|
||||||
|
}
|
||||||
|
for sum > 0xffff {
|
||||||
|
sum = (sum >> 16) + (sum & 0xffff)
|
||||||
|
}
|
||||||
|
return ^uint16(sum)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ipv4HeaderChecksumValid reports whether the stored checksum matches a fresh
|
||||||
|
// computation over the header.
|
||||||
|
func ipv4HeaderChecksumValid(hdr []byte) bool {
|
||||||
|
return binary.BigEndian.Uint16(hdr[10:12]) == ipv4HeaderChecksum(hdr)
|
||||||
|
}
|
||||||
+20
-2
@@ -255,15 +255,23 @@ tun:
|
|||||||
mtu: 1300
|
mtu: 1300
|
||||||
|
|
||||||
# Linux only. pin_threads pins each tun reader/encrypt OS thread to a single CPU. This keeps every goroutine's
|
# Linux only. pin_threads pins each tun reader/encrypt OS thread to a single CPU. This keeps every goroutine's
|
||||||
# sends flowing through one XPS-selected NIC TX ring, so packets within a flow stay ordered on the wire
|
# batched sends flowing through one XPS-selected NIC TX ring, so packets within a flow stay ordered on the wire
|
||||||
# instead of being sprayed across multiple TX rings and reordered. Not reloadable.
|
# instead of being sprayed across multiple TX rings and reordered. Not reloadable.
|
||||||
|
#
|
||||||
|
# When cpu_affinity is unset, nebula picks CPUs that do NOT service any physical NIC's interrupts (read from
|
||||||
|
# /sys/class/net/*/device/msi_irqs and /proc/irq/*/effective_affinity_list): an encrypt thread pinned onto a core
|
||||||
|
# that also runs NAPI for a NIC RX queue fights the softirq for the core and collapses throughput for flows hashed
|
||||||
|
# to that queue. If the NIC's vectors blanket every allowed CPU (many drivers default to one queue per core) the
|
||||||
|
# avoidance logs and falls back to the old spread; narrow the NIC's queue/IRQ spread (e.g. `ethtool -X <dev>
|
||||||
|
# equal N`) or set cpu_affinity explicitly to benefit.
|
||||||
#pin_threads: true
|
#pin_threads: true
|
||||||
|
|
||||||
# Linux only. cpu_affinity overrides which CPUs the tun reader threads pin to: a list of CPU IDs, one per routine
|
# Linux only. cpu_affinity overrides which CPUs the tun reader threads pin to: a list of CPU IDs, one per routine
|
||||||
# (see the top-level `routines` setting). Lists shorter than `routines` are modulo-cycled across the queues; extra
|
# (see the top-level `routines` setting). Lists shorter than `routines` are modulo-cycled across the queues; extra
|
||||||
# entries are ignored. IDs must be within the process's allowed CPU set, so this respects taskset / cgroup cpusets;
|
# entries are ignored. IDs must be within the process's allowed CPU set, so this respects taskset / cgroup cpusets;
|
||||||
# a non-integer or not-allowed entry disables the override and falls back to spreading queues across the allowed
|
# a non-integer or not-allowed entry disables the override and falls back to spreading queues across the allowed
|
||||||
# CPUs. Only meaningful while pin_threads is true. Not reloadable.
|
# CPUs. Setting this disables the automatic NIC-IRQ avoidance described under pin_threads — prefer CPUs that don't
|
||||||
|
# service your underlay NIC's RX queue IRQs. Only meaningful while pin_threads is true. Not reloadable.
|
||||||
#cpu_affinity:
|
#cpu_affinity:
|
||||||
# - 2
|
# - 2
|
||||||
# - 4
|
# - 4
|
||||||
@@ -404,6 +412,16 @@ logging:
|
|||||||
# This setting is reloadable
|
# This setting is reloadable
|
||||||
#inactivity_timeout: 10m
|
#inactivity_timeout: 10m
|
||||||
|
|
||||||
|
# ecn (default true) propagates ECN (Explicit Congestion Notification) across the tunnel per RFC 6040: the inner
|
||||||
|
# packet's ECN codepoint is copied onto the outer carrier header on encapsulation, and an outer CE ("congestion
|
||||||
|
# experienced") mark is folded back into the inner header on decapsulation. On linux it additionally stamps
|
||||||
|
# RTAX_FEATURE_ECN on the routes nebula installs, so the kernel actively negotiates ECN for connections to mesh
|
||||||
|
# prefixes. Disable this only when an underlay middlebox mangles or clears ECN bits unpredictably.
|
||||||
|
# This setting is reloadable, BUT flipping it at runtime only updates the datapath (the inner<->outer copy/combine).
|
||||||
|
# The RTAX_FEATURE_ECN flag on already-installed routes is NOT revisited on reload, so nebula must be restarted for
|
||||||
|
# the route half of this setting to take effect.
|
||||||
|
#ecn: true
|
||||||
|
|
||||||
# Nebula security group configuration
|
# Nebula security group configuration
|
||||||
firewall:
|
firewall:
|
||||||
# Action to take when a packet is not allowed by the firewall rules.
|
# Action to take when a packet is not allowed by the firewall rules.
|
||||||
|
|||||||
+4
-2
@@ -5,6 +5,8 @@ import (
|
|||||||
"log/slog"
|
"log/slog"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/logging"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ConntrackCache is used as a local routine cache to know if a given flow
|
// ConntrackCache is used as a local routine cache to know if a given flow
|
||||||
@@ -56,8 +58,8 @@ func (c *ConntrackCacheTicker) Get() ConntrackCache {
|
|||||||
if tick := c.cacheTick.Load(); tick != c.cacheV {
|
if tick := c.cacheTick.Load(); tick != c.cacheV {
|
||||||
c.cacheV = tick
|
c.cacheV = tick
|
||||||
if ll := len(c.cache); ll > 0 {
|
if ll := len(c.cache); ll > 0 {
|
||||||
if c.l.Enabled(context.Background(), slog.LevelDebug) {
|
if c.l.Enabled(context.Background(), logging.LevelTrace) {
|
||||||
c.l.Debug("resetting conntrack cache", "len", ll)
|
c.l.Log(context.Background(), logging.LevelTrace, "resetting conntrack cache", "len", ll)
|
||||||
}
|
}
|
||||||
c.cache = make(ConntrackCache, ll)
|
c.cache = make(ConntrackCache, ll)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/logging"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
)
|
)
|
||||||
@@ -30,27 +31,27 @@ func newFixedTicker(t *testing.T, l *slog.Logger, cacheLen int) *ConntrackCacheT
|
|||||||
|
|
||||||
func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) {
|
func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 3)
|
c := newFixedTicker(t, l, 3)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|
||||||
assert.Equal(t, "level=DEBUG msg=\"resetting conntrack cache\" len=3\n", buf.String())
|
assert.Equal(t, "level=DEBUG-4 msg=\"resetting conntrack cache\" len=3\n", buf.String())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) {
|
func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewJSONLoggerWithOutput(buf, slog.LevelDebug)
|
l := test.NewJSONLoggerWithOutput(buf, logging.LevelTrace)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 2)
|
c := newFixedTicker(t, l, 2)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|
||||||
assert.JSONEq(t, `{"level":"DEBUG","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
|
assert.JSONEq(t, `{"level":"DEBUG-4","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestConntrackCacheTicker_Get_QuietBelowDebug(t *testing.T) {
|
func TestConntrackCacheTicker_Get_QuietBelowTrace(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelInfo)
|
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 5)
|
c := newFixedTicker(t, l, 5)
|
||||||
c.Get()
|
c.Get()
|
||||||
@@ -60,7 +61,7 @@ func TestConntrackCacheTicker_Get_QuietBelowDebug(t *testing.T) {
|
|||||||
|
|
||||||
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
|
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 0)
|
c := newFixedTicker(t, l, 0)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|||||||
@@ -967,6 +967,7 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
|||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
out := make([]byte, mtu)
|
out := make([]byte, mtu)
|
||||||
for _, cp := range hh.packetStore {
|
for _, cp := range hh.packetStore {
|
||||||
|
//todo use a sendbatcher
|
||||||
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
||||||
}
|
}
|
||||||
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
@@ -9,10 +10,24 @@ import (
|
|||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/iputil"
|
"github.com/slackhq/nebula/iputil"
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packet, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
|
// borrowed: pkt.Bytes is owned by the originating tio.Queue and is
|
||||||
|
// only valid until the next Read on that queue. Every consumer below
|
||||||
|
// (parse, self-forward, handshake cache, sendInsideMessage) reads it
|
||||||
|
// synchronously; do not retain pkt outside this call. If a future
|
||||||
|
// caller needs to keep the packet, use pkt.Clone() to detach it from
|
||||||
|
// the borrow.
|
||||||
|
//
|
||||||
|
// pkt.Bytes is either one IP datagram (GSO zero) or a TSO/USO
|
||||||
|
// superpacket. In both cases the L3+L4 headers at the start describe
|
||||||
|
// the same 5-tuple every segment will share, so a single newPacket /
|
||||||
|
// firewall check covers the whole superpacket.
|
||||||
|
packet := pkt.Bytes
|
||||||
err := newPacket(packet, false, fwPacket)
|
err := newPacket(packet, false, fwPacket)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
@@ -37,7 +52,14 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
||||||
// TUN device.
|
// TUN device.
|
||||||
if immediatelyForwardToSelf {
|
if immediatelyForwardToSelf {
|
||||||
_, err := f.queues[q].Write(packet)
|
// Write copies into the kernel queue synchronously, so seg's lifetime ends at return.
|
||||||
|
// A self-forwarded superpacket would be re-handed to the
|
||||||
|
// kernel as one giant blob; segment first so the loopback
|
||||||
|
// path sees one IP datagram per Write.
|
||||||
|
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
||||||
|
_, werr := f.queues[q].Write(seg)
|
||||||
|
return werr
|
||||||
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Failed to forward to tun", "error", err)
|
f.l.Error("Failed to forward to tun", "error", err)
|
||||||
}
|
}
|
||||||
@@ -53,11 +75,23 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
}
|
}
|
||||||
|
|
||||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
||||||
hh.cachePacket(f.l, header.Message, 0, packet, f.sendMessageNow, f.cachedPacketMetrics)
|
// borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt
|
||||||
|
// bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket),
|
||||||
|
// so retaining segments past the loop is safe.
|
||||||
|
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
||||||
|
hh.cachePacket(f.l, header.Message, 0, seg, f.sendMessageNow, f.cachedPacketMetrics)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil && f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("Failed to segment superpacket for handshake cache",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddr", fwPacket.RemoteAddr,
|
||||||
|
)
|
||||||
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
||||||
"vpnAddr", fwPacket.RemoteAddr,
|
"vpnAddr", fwPacket.RemoteAddr,
|
||||||
@@ -73,10 +107,9 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
|
|
||||||
dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
|
dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
|
||||||
if dropReason == nil {
|
if dropReason == nil {
|
||||||
f.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q)
|
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch, rejectBuf, q)
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
||||||
"fwPacket", fwPacket,
|
"fwPacket", fwPacket,
|
||||||
@@ -86,6 +119,150 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, seg, scratch, nb []byte) []byte {
|
||||||
|
if noiseutil.EncryptLockNeeded {
|
||||||
|
ci.writeLock.Lock()
|
||||||
|
}
|
||||||
|
c := ci.messageCounter.Add(1)
|
||||||
|
|
||||||
|
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
|
||||||
|
f.connectionManager.Out(hostinfo)
|
||||||
|
|
||||||
|
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.TxBatcher, rejectBuf []byte, q int) {
|
||||||
|
ci := hostinfo.ConnectionState
|
||||||
|
if ci.eKey == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
remote := hostinfo.GetRemote()
|
||||||
|
ecnEnabled := f.ecnEnabled.Load()
|
||||||
|
if hostinfo.lastRebindCount != f.rebindCount {
|
||||||
|
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
|
||||||
|
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
|
||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
var ecn byte
|
||||||
|
if ecnEnabled {
|
||||||
|
ecn = innerECN(seg)
|
||||||
|
}
|
||||||
|
sendBatch.Commit(toSend, relayHostInfo.GetRemote(), ecn)
|
||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
var ecn byte
|
||||||
|
if ecnEnabled {
|
||||||
|
ecn = innerECN(seg)
|
||||||
|
}
|
||||||
|
sendBatch.Commit(out, remote, ecn)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(f.l).Error("Failed to segment superpacket for send",
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// innerECN returns the 2-bit IP-level ECN codepoint of an inner IPv4 or IPv6
|
||||||
|
// packet, or 0 if pkt is too short or its IP version is unrecognized. Used at
|
||||||
|
// encap to copy the inner codepoint onto the outer carrier per RFC 6040.
|
||||||
|
func innerECN(pkt []byte) byte {
|
||||||
|
if len(pkt) < 2 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
switch pkt[0] >> 4 {
|
||||||
|
case 4:
|
||||||
|
return pkt[1] & 0x03
|
||||||
|
case 6:
|
||||||
|
return (pkt[1] >> 4) & 0x03
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||||
if !f.firewall.OutboundSendReject {
|
if !f.firewall.OutboundSendReject {
|
||||||
return
|
return
|
||||||
@@ -275,21 +452,13 @@ func (f *Interface) sendTo(t header.MessageType, st header.MessageSubType, ci *C
|
|||||||
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
|
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
func (f *Interface) prepareSendVia(via *HostInfo,
|
||||||
// to the payload for the ultimate target host, making this a useful method for sending
|
|
||||||
// handshake messages to peers through relay tunnels.
|
|
||||||
// via is the HostInfo through which the message is relayed.
|
|
||||||
// ad is the plaintext data to authenticate, but not encrypt
|
|
||||||
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
|
||||||
// out is a buffer used to store the result of the Encrypt operation
|
|
||||||
// q indicates which writer to use to send the packet.
|
|
||||||
func (f *Interface) SendVia(via *HostInfo,
|
|
||||||
relay *Relay,
|
relay *Relay,
|
||||||
ad,
|
ad,
|
||||||
nb,
|
nb,
|
||||||
out []byte,
|
out []byte,
|
||||||
nocopy bool,
|
nocopy bool,
|
||||||
) {
|
) ([]byte, error) {
|
||||||
if noiseutil.EncryptLockNeeded {
|
if noiseutil.EncryptLockNeeded {
|
||||||
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
||||||
via.ConnectionState.writeLock.Lock()
|
via.ConnectionState.writeLock.Lock()
|
||||||
@@ -311,7 +480,7 @@ func (f *Interface) SendVia(via *HostInfo,
|
|||||||
"headerLen", len(out),
|
"headerLen", len(out),
|
||||||
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
||||||
)
|
)
|
||||||
return
|
return nil, io.ErrShortBuffer
|
||||||
}
|
}
|
||||||
|
|
||||||
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
||||||
@@ -331,13 +500,37 @@ func (f *Interface) SendVia(via *HostInfo,
|
|||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
f.connectionManager.RelayUsed(relay.LocalIndex)
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
||||||
|
// to the payload for the ultimate target host, making this a useful method for sending
|
||||||
|
// handshake messages to peers through relay tunnels.
|
||||||
|
// via is the HostInfo through which the message is relayed.
|
||||||
|
// ad is the plaintext data to authenticate, but not encrypt
|
||||||
|
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
||||||
|
// out is a buffer used to store the result of the Encrypt operation
|
||||||
|
// q indicates which writer to use to send the packet.
|
||||||
|
func (f *Interface) SendVia(via *HostInfo,
|
||||||
|
relay *Relay,
|
||||||
|
ad,
|
||||||
|
nb,
|
||||||
|
out []byte,
|
||||||
|
nocopy bool,
|
||||||
|
) {
|
||||||
|
toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy)
|
||||||
|
if err != nil {
|
||||||
|
// already logged by prepareSendVia
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
err = f.writers[0].WriteTo(out, via.GetRemote())
|
|
||||||
|
err = f.writers[0].WriteTo(toSend, via.GetRemote())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
||||||
}
|
}
|
||||||
f.connectionManager.RelayUsed(relay.LocalIndex)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
||||||
|
|||||||
+77
-15
@@ -14,15 +14,16 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
|
"github.com/slackhq/nebula/util"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
"github.com/slackhq/nebula/util"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const mtu = 9001
|
const mtu = 9001
|
||||||
@@ -59,7 +60,7 @@ type InterfaceConfig struct {
|
|||||||
CpuAffinity []int
|
CpuAffinity []int
|
||||||
// PinThreads controls whether each TUN reader OS thread is pinned to a
|
// PinThreads controls whether each TUN reader OS thread is pinned to a
|
||||||
// single CPU (via tun.pin_threads, default true). Pinning keeps each
|
// single CPU (via tun.pin_threads, default true). Pinning keeps each
|
||||||
// goroutine's UDP sends on one XPS-selected NIC TX ring so per-flow
|
// goroutine's sendmmsg on one XPS-selected NIC TX ring so per-flow
|
||||||
// packets stay ordered on the wire.
|
// packets stay ordered on the wire.
|
||||||
PinThreads bool
|
PinThreads bool
|
||||||
|
|
||||||
@@ -95,7 +96,12 @@ type Interface struct {
|
|||||||
// pinThreads controls whether listenIn pins each TUN reader OS thread to
|
// pinThreads controls whether listenIn pins each TUN reader OS thread to
|
||||||
// a CPU at all (tun.pin_threads, default true). When false, threads are
|
// a CPU at all (tun.pin_threads, default true). When false, threads are
|
||||||
// left free to migrate as on stock nebula.
|
// left free to migrate as on stock nebula.
|
||||||
pinThreads bool
|
pinThreads bool
|
||||||
|
// ecnEnabled gates RFC 6040 underlay ECN propagation. When true,
|
||||||
|
// inside.go copies the inner ECN onto the outer carrier on encap and
|
||||||
|
// decryptToTun folds outer CE into the inner header on decap. Toggle
|
||||||
|
// via tunnels.ecn (default true).
|
||||||
|
ecnEnabled atomic.Bool
|
||||||
relayManager *relayManager
|
relayManager *relayManager
|
||||||
|
|
||||||
tryPromoteEvery atomic.Uint32
|
tryPromoteEvery atomic.Uint32
|
||||||
@@ -114,7 +120,11 @@ type Interface struct {
|
|||||||
ctx context.Context
|
ctx context.Context
|
||||||
writers []udp.Conn
|
writers []udp.Conn
|
||||||
queues []tio.Queue
|
queues []tio.Queue
|
||||||
wg sync.WaitGroup
|
// batchers is one per tun queue, wrapping queues[i].
|
||||||
|
// decryptToTun sends plaintext into the batch.RxBatcher;
|
||||||
|
// listenOut calls its Flush at the end of each UDP recvmmsg batch.
|
||||||
|
batchers []batch.RxBatcher
|
||||||
|
wg sync.WaitGroup
|
||||||
|
|
||||||
// fatalErr holds the first unexpected reader error that caused shutdown.
|
// fatalErr holds the first unexpected reader error that caused shutdown.
|
||||||
// nil means "no fatal error" (yet)
|
// nil means "no fatal error" (yet)
|
||||||
@@ -212,6 +222,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
routines: c.routines,
|
routines: c.routines,
|
||||||
version: c.version,
|
version: c.version,
|
||||||
writers: make([]udp.Conn, c.routines),
|
writers: make([]udp.Conn, c.routines),
|
||||||
|
batchers: make([]batch.RxBatcher, c.routines),
|
||||||
myVpnNetworks: cs.myVpnNetworks,
|
myVpnNetworks: cs.myVpnNetworks,
|
||||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||||
myVpnAddrs: cs.myVpnAddrs,
|
myVpnAddrs: cs.myVpnAddrs,
|
||||||
@@ -285,6 +296,21 @@ func (f *Interface) activate() error {
|
|||||||
|
|
||||||
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
||||||
|
|
||||||
|
for i := range f.queues {
|
||||||
|
caps := tio.QueueCapabilities(f.queues[i])
|
||||||
|
if caps.TSO || caps.USO {
|
||||||
|
// Multi-lane: TCP gets coalesced when TSO is on, UDP when USO
|
||||||
|
// is on, everything else (and either lane disabled) falls
|
||||||
|
// through to passthrough so non-IP / non-TCP-UDP traffic still
|
||||||
|
// reaches the TUN.
|
||||||
|
arena := batch.NewArena(batch.DefaultMultiArenaCap)
|
||||||
|
f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l, arena, caps.TSO, caps.USO)
|
||||||
|
} else {
|
||||||
|
arena := batch.NewArena(batch.DefaultPassthroughArenaCap)
|
||||||
|
f.batchers[i] = batch.NewPassthrough(f.queues[i], arena.Reserve, arena.Reset)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// On error the caller owns the cleanup, Control.Start cancels the service context
|
// 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
|
// before releasing our resources so a waiter never observes a live context
|
||||||
if err = f.inside.Activate(); err != nil {
|
if err = f.inside.Activate(); err != nil {
|
||||||
@@ -340,14 +366,22 @@ func (f *Interface) listenOut(i int) {
|
|||||||
|
|
||||||
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||||
lhh := f.lightHouse.NewRequestHandler()
|
lhh := f.lightHouse.NewRequestHandler()
|
||||||
plaintext := make([]byte, udp.MTU)
|
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
fwPacket := &firewall.Packet{}
|
fwPacket := &firewall.Packet{}
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
|
listener := func(fromUdpAddr netip.AddrPort, payload []byte, meta udp.RxMeta) {
|
||||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
plaintext := f.batchers[i].Reserve(len(payload))
|
||||||
})
|
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get(), meta)
|
||||||
|
}
|
||||||
|
|
||||||
|
flusher := func() {
|
||||||
|
if err := f.batchers[i].Flush(); err != nil {
|
||||||
|
f.l.Error("Failed to flush tun coalescer", "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err := li.ListenOut(listener, flusher)
|
||||||
|
|
||||||
// An error after teardown began is shutdown noise, the closed flag covers resources
|
// An error after teardown began is shutdown noise, the closed flag covers resources
|
||||||
// Close releases itself and the cancelled ctx covers ones torn down by their owners
|
// Close releases itself and the cancelled ctx covers ones torn down by their owners
|
||||||
@@ -361,9 +395,8 @@ func (f *Interface) listenOut(i int) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) listenIn(queue tio.Queue, i int) {
|
func (f *Interface) listenIn(queue tio.Queue, i int) {
|
||||||
// Pinning this thread (and goroutine) to a single CPU keeps every UDP send from this goroutine going through
|
// Pinning this thread (and goroutine) to a single CPU keeps every sendmmsg from this goroutine going through the
|
||||||
// the same TX ring on the nic (XPS selects the ring by CPU), so the wire sees per-flow order. Skip entirely
|
// same TX ring on the nic, so the wire sees per-flow order. Skip entirely when tun.pin_threads is false.
|
||||||
// when tun.pin_threads is false.
|
|
||||||
if f.pinThreads {
|
if f.pinThreads {
|
||||||
var cpu int
|
var cpu int
|
||||||
if n := len(f.cpuAffinity); n > 0 {
|
if n := len(f.cpuAffinity); n > 0 {
|
||||||
@@ -383,7 +416,9 @@ func (f *Interface) listenIn(queue tio.Queue, i int) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
out := make([]byte, mtu)
|
rejectBuf := make([]byte, mtu)
|
||||||
|
arenaSize := batch.SendBatchCap * (udp.MTU + 32)
|
||||||
|
sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, arenaSize)
|
||||||
fwPacket := &firewall.Packet{}
|
fwPacket := &firewall.Packet{}
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
@@ -401,9 +436,18 @@ func (f *Interface) listenIn(queue tio.Queue, i int) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
for _, pkt := range pkts {
|
for _, pkt := range pkts {
|
||||||
// borrowed: pkt.Bytes is owned by the queue and only valid until
|
f.consumeInsidePacket(pkt, fwPacket, nb, sb, rejectBuf, i, conntrackCache.Get())
|
||||||
// the next Read; consumeInsidePacket reads it synchronously.
|
// Flush incrementally once a full sendmmsg batch has
|
||||||
f.consumeInsidePacket(pkt.Bytes, fwPacket, nb, out, i, conntrackCache.Get())
|
// 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 {
|
||||||
|
if err := sb.Flush(); err != nil {
|
||||||
|
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := sb.Flush(); err != nil {
|
||||||
|
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -416,6 +460,7 @@ func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
|||||||
c.RegisterReloadCallback(f.reloadAcceptRecvError)
|
c.RegisterReloadCallback(f.reloadAcceptRecvError)
|
||||||
c.RegisterReloadCallback(f.reloadDisconnectInvalid)
|
c.RegisterReloadCallback(f.reloadDisconnectInvalid)
|
||||||
c.RegisterReloadCallback(f.reloadMisc)
|
c.RegisterReloadCallback(f.reloadMisc)
|
||||||
|
c.RegisterReloadCallback(f.reloadEcn)
|
||||||
|
|
||||||
for _, udpConn := range f.writers {
|
for _, udpConn := range f.writers {
|
||||||
c.RegisterReloadCallback(udpConn.ReloadConfig)
|
c.RegisterReloadCallback(udpConn.ReloadConfig)
|
||||||
@@ -548,6 +593,23 @@ func (f *Interface) reloadMisc(c *config.C) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// reloadEcn syncs Interface.ecnEnabled with the tunnels.ecn config knob.
|
||||||
|
// Default is enabled (RFC 6040 normal mode); set false on the rare path
|
||||||
|
// where an underlay middlebox rewrites or drops ECN bits unpredictably.
|
||||||
|
func (f *Interface) reloadEcn(c *config.C) {
|
||||||
|
initial := c.InitialLoad()
|
||||||
|
if initial || c.HasChanged("tunnels.ecn") {
|
||||||
|
v := c.GetBool("tunnels.ecn", true)
|
||||||
|
changed := f.ecnEnabled.Swap(v) != v
|
||||||
|
if !initial {
|
||||||
|
f.l.Info("tunnels.ecn changed", "enabled", v)
|
||||||
|
if changed {
|
||||||
|
f.l.Warn("tunnels.ecn datapath toggled, but route-level ECN negotiation (RTAX_FEATURE_ECN) retains its previous state until nebula is restarted", "enabled", v)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
||||||
ticker := time.NewTicker(i)
|
ticker := time.NewTicker(i)
|
||||||
defer ticker.Stop()
|
defer ticker.Stop()
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package iputil
|
package iputil
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"net"
|
"net"
|
||||||
"testing"
|
"testing"
|
||||||
@@ -179,6 +180,46 @@ func Test_CreateRejectPacket_NoICMPError(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Test_CreateRejectPacket_RespectsCap ensures it is impossible for
|
||||||
|
// an oversized ICMPv6 reject to overwrite the neighbor segment's bytes.
|
||||||
|
func Test_CreateRejectPacket_RespectsCap(t *testing.T) {
|
||||||
|
src := net.ParseIP("fd00::1")
|
||||||
|
dst := net.ParseIP("fd00::2")
|
||||||
|
|
||||||
|
// Inner IPv6 UDP packet. An ICMPv6 reject copies the whole inner packet
|
||||||
|
// plus a 48-byte header (40 IPv6 + 8 ICMPv6), so it needs 48 more bytes
|
||||||
|
// than the inner packet length.
|
||||||
|
inner := makeIPv6Packet(src, dst, 17, make([]byte, 20))
|
||||||
|
|
||||||
|
// The ciphertext scratch reused as the reject buffer is the received
|
||||||
|
// datagram: 16-byte Nebula header + inner + 16-byte AEAD tag. That is only
|
||||||
|
// 32 bytes of slack, so a full ICMPv6 reject overruns it by 16 bytes.
|
||||||
|
const nebulaOverhead = 32
|
||||||
|
segLen := len(inner) + nebulaOverhead
|
||||||
|
|
||||||
|
// Shared backing row laid out as [segment][neighbor's 16-byte Nebula header].
|
||||||
|
const neighborHdr = 16
|
||||||
|
sentinel := bytes.Repeat([]byte{0xAB}, neighborHdr)
|
||||||
|
|
||||||
|
// Uncapped: the slice's capacity reaches into the neighbor, reproducing
|
||||||
|
// the overrun that silently drops the neighbor packet.
|
||||||
|
backing := make([]byte, segLen+neighborHdr)
|
||||||
|
copy(backing[segLen:], sentinel)
|
||||||
|
reject := CreateRejectPacket(inner, backing[:segLen])
|
||||||
|
assert.NotNil(t, reject, "uncapped buffer reaches into the neighbor, so the reject is built")
|
||||||
|
assert.NotEqual(t, sentinel, backing[segLen:segLen+neighborHdr],
|
||||||
|
"without the cap the oversized reject overruns into the neighbor segment")
|
||||||
|
|
||||||
|
// Capped (the fix): cap==len, so the builder cannot exceed the segment. The
|
||||||
|
// reject does not fit, so it is refused rather than corrupting the neighbor.
|
||||||
|
backing = make([]byte, segLen+neighborHdr)
|
||||||
|
copy(backing[segLen:], sentinel)
|
||||||
|
reject = CreateRejectPacket(inner, backing[:segLen:segLen])
|
||||||
|
assert.Nil(t, reject, "capped segment is 16 bytes too small for a full ICMPv6 reject, so it is refused")
|
||||||
|
assert.Equal(t, sentinel, backing[segLen:segLen+neighborHdr],
|
||||||
|
"capped segment must leave the neighbor untouched")
|
||||||
|
}
|
||||||
|
|
||||||
func makeIPv6Packet(src, dst net.IP, nextHeader uint8, payload []byte) []byte {
|
func makeIPv6Packet(src, dst net.IP, nextHeader uint8, payload []byte) []byte {
|
||||||
b := make([]byte, ipv6.HeaderLen+len(payload))
|
b := make([]byte, ipv6.HeaderLen+len(payload))
|
||||||
b[0] = ipv6.Version << 4
|
b[0] = ipv6.Version << 4
|
||||||
|
|||||||
@@ -34,6 +34,9 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
buildVersion = moduleVersion()
|
buildVersion = moduleVersion()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Debug builds (-tags debug) serve pprof on :6060; a no-op otherwise.
|
||||||
|
startPprofServer(ctx, l)
|
||||||
|
|
||||||
// Print the config if in test, the exit comes later
|
// Print the config if in test, the exit comes later
|
||||||
if configTest {
|
if configTest {
|
||||||
b, err := yaml.Marshal(c.Settings)
|
b, err := yaml.Marshal(c.Settings)
|
||||||
@@ -211,6 +214,12 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
l.Warn("Failed to start DNS responder", "error", err)
|
l.Warn("Failed to start DNS responder", "error", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pinThreads := c.GetBool("tun.pin_threads", true)
|
||||||
|
cpuAffinity := parseCpuAffinity(c, l, routines)
|
||||||
|
if pinThreads && len(cpuAffinity) == 0 && !configTest {
|
||||||
|
cpuAffinity = defaultCPUAffinityAvoidingIRQs(l, routines)
|
||||||
|
}
|
||||||
|
|
||||||
ifConfig := &InterfaceConfig{
|
ifConfig := &InterfaceConfig{
|
||||||
HostMap: hostMap,
|
HostMap: hostMap,
|
||||||
Inside: tun,
|
Inside: tun,
|
||||||
@@ -232,8 +241,8 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
relayManager: NewRelayManager(ctx, l, hostMap, c),
|
relayManager: NewRelayManager(ctx, l, hostMap, c),
|
||||||
punchy: punchy,
|
punchy: punchy,
|
||||||
ConntrackCacheTimeout: conntrackCacheTimeout,
|
ConntrackCacheTimeout: conntrackCacheTimeout,
|
||||||
CpuAffinity: parseCpuAffinity(c, l, routines),
|
CpuAffinity: cpuAffinity,
|
||||||
PinThreads: c.GetBool("tun.pin_threads", true),
|
PinThreads: pinThreads,
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -251,6 +260,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
ifce.reloadDisconnectInvalid(c)
|
ifce.reloadDisconnectInvalid(c)
|
||||||
ifce.reloadSendRecvError(c)
|
ifce.reloadSendRecvError(c)
|
||||||
ifce.reloadAcceptRecvError(c)
|
ifce.reloadAcceptRecvError(c)
|
||||||
|
ifce.reloadEcn(c)
|
||||||
|
|
||||||
handshakeManager.f = ifce
|
handshakeManager.f = ifce
|
||||||
go handshakeManager.Run(ctx)
|
go handshakeManager.Run(ctx)
|
||||||
@@ -349,6 +359,57 @@ func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int {
|
|||||||
return cpus
|
return cpus
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// defaultCPUAffinityAvoidingIRQs picks the default pin set for the tun
|
||||||
|
// readers when tun.cpu_affinity is unset: allowed CPUs that do NOT service
|
||||||
|
// any physical NIC's interrupts. The stock allowed[i] spread pins the
|
||||||
|
// encrypt threads onto exactly the cores most drivers affine their first RX
|
||||||
|
// queue IRQs to, so whenever a flow's RSS queue fires on a core hosting a
|
||||||
|
// tun reader, NAPI and encrypt fight for the core and per-flow throughput
|
||||||
|
// drops (measured: REV 8.4 vs 10.2 Gbps on the same hardware, 2026-07-14).
|
||||||
|
//
|
||||||
|
// Returns nil — keeping the old allowed[i] fallback in listenIn — when IRQ
|
||||||
|
// info is unavailable or when there aren't enough IRQ-free CPUs to give
|
||||||
|
// every routine its own core: silently doubling readers up on fewer cores
|
||||||
|
// is worse than the occasional IRQ collision. NICs whose vectors blanket
|
||||||
|
// every CPU (e.g. mlx5 defaults to one queue per core) make avoidance
|
||||||
|
// impossible; narrowing the NIC's spread (ethtool -X <dev> equal N, or
|
||||||
|
// /proc/irq/*/smp_affinity) or setting tun.cpu_affinity explicitly makes it
|
||||||
|
// effective.
|
||||||
|
func defaultCPUAffinityAvoidingIRQs(l *slog.Logger, routines int) []int {
|
||||||
|
irq, err := util.NICIRQCPUs()
|
||||||
|
if err != nil || len(irq) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
allowed, err := util.AllowedCPUs()
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
cpus := chooseIRQFreeCPUs(allowed, irq, routines)
|
||||||
|
if cpus == nil {
|
||||||
|
l.Info("not enough CPUs are free of NIC IRQs to give every tun reader its own; using the default spread",
|
||||||
|
"routines", routines, "allowed", len(allowed), "irqCPUs", len(irq))
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
l.Info("pinning tun readers to CPUs clear of NIC IRQs", "cpus", cpus)
|
||||||
|
return cpus
|
||||||
|
}
|
||||||
|
|
||||||
|
// chooseIRQFreeCPUs returns the first `routines` allowed CPUs not present in
|
||||||
|
// irq, or nil if fewer than `routines` qualify.
|
||||||
|
func chooseIRQFreeCPUs(allowed []int, irq map[int]bool, routines int) []int {
|
||||||
|
free := make([]int, 0, routines)
|
||||||
|
for _, cpu := range allowed {
|
||||||
|
if irq[cpu] {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
free = append(free, cpu)
|
||||||
|
if len(free) == routines {
|
||||||
|
return free
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func moduleVersion() string {
|
func moduleVersion() string {
|
||||||
info, ok := debug.ReadBuildInfo()
|
info, ok := debug.ReadBuildInfo()
|
||||||
if !ok {
|
if !ok {
|
||||||
|
|||||||
@@ -9,6 +9,26 @@ import (
|
|||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
func TestChooseIRQFreeCPUs(t *testing.T) {
|
||||||
|
irq := map[int]bool{0: true, 1: true, 2: true, 3: true}
|
||||||
|
|
||||||
|
// Plenty of IRQ-free CPUs: take the first `routines` of them in order.
|
||||||
|
assert.Equal(t, []int{4, 5}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 2))
|
||||||
|
|
||||||
|
// Exactly enough.
|
||||||
|
assert.Equal(t, []int{4, 5, 6}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 3))
|
||||||
|
|
||||||
|
// Not enough IRQ-free CPUs: nil, caller keeps the old default rather
|
||||||
|
// than doubling readers up on shared cores.
|
||||||
|
assert.Nil(t, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4}, irq, 2))
|
||||||
|
|
||||||
|
// No IRQ info at all behaves like a plain prefix of allowed.
|
||||||
|
assert.Equal(t, []int{0, 1}, chooseIRQFreeCPUs([]int{0, 1, 2}, map[int]bool{}, 2))
|
||||||
|
|
||||||
|
// Non-contiguous allowed set (cgroup cpuset) with holes.
|
||||||
|
assert.Equal(t, []int{9, 12}, chooseIRQFreeCPUs([]int{1, 3, 9, 12}, map[int]bool{1: true, 3: true}, 2))
|
||||||
|
}
|
||||||
|
|
||||||
func TestParseCpuAffinity(t *testing.T) {
|
func TestParseCpuAffinity(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
|
|
||||||
|
|||||||
+83
-11
@@ -13,6 +13,7 @@ import (
|
|||||||
|
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -22,7 +23,7 @@ const (
|
|||||||
|
|
||||||
var ErrOutOfWindow = errors.New("out of window packet")
|
var ErrOutOfWindow = errors.New("out of window packet")
|
||||||
|
|
||||||
func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
|
||||||
err := h.Parse(packet)
|
err := h.Parse(packet)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
||||||
@@ -110,8 +111,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
|
|
||||||
// Relay packets are special
|
// Relay packets are special
|
||||||
if isMessageRelay {
|
if isMessageRelay {
|
||||||
f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, lhf, nb, q, localCache)
|
f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, lhf, nb, q, localCache, meta)
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -135,7 +135,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
case header.Message:
|
case header.Message:
|
||||||
switch h.Subtype {
|
switch h.Subtype {
|
||||||
case header.MessageNone:
|
case header.MessageNone:
|
||||||
f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, nb, q, localCache)
|
f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, nb, q, localCache, meta)
|
||||||
default:
|
default:
|
||||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
|
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
|
||||||
return
|
return
|
||||||
@@ -169,7 +169,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
|
||||||
// The entire body is sent as AD, not encrypted.
|
// The 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 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
|
// 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
|
||||||
@@ -218,7 +218,7 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
|||||||
relay: relay,
|
relay: relay,
|
||||||
IsRelayed: true,
|
IsRelayed: true,
|
||||||
}
|
}
|
||||||
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache)
|
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache, meta)
|
||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
// Find the target HostInfo relay object
|
// Find the target HostInfo relay object
|
||||||
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
||||||
@@ -236,7 +236,7 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
|||||||
switch targetRelay.Type {
|
switch targetRelay.Type {
|
||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
// Forward this packet through the relay tunnel
|
// Forward this packet through the relay tunnel
|
||||||
// Find the target HostInfo
|
// Find the target HostInfo //todo it would potentially be nice to batch these
|
||||||
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
||||||
case TerminalType:
|
case TerminalType:
|
||||||
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
||||||
@@ -518,7 +518,77 @@ func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, out []byte, packet []
|
|||||||
return out, nil
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) {
|
// 2-bit IP-level ECN codepoints (lower bits of IPv4 ToS / IPv6 TC).
|
||||||
|
const (
|
||||||
|
ecnNotECT = 0x00
|
||||||
|
ecnECT1 = 0x01
|
||||||
|
ecnECT0 = 0x02
|
||||||
|
ecnCE = 0x03
|
||||||
|
)
|
||||||
|
|
||||||
|
// applyOuterECN folds an outer CE mark from the underlay into the inner
|
||||||
|
// IP header per RFC 6040 normal mode. It mutates pkt[1] in place. Other
|
||||||
|
// codepoints are advisory only and leave the inner unchanged.
|
||||||
|
//
|
||||||
|
// Merge cases (outer × inner → action):
|
||||||
|
//
|
||||||
|
// outer != CE : no-op (inner is authoritative)
|
||||||
|
// outer == CE, inner Not-ECT : log; cannot propagate to a non-ECN host
|
||||||
|
// outer == CE, inner ECT/CE : rewrite inner ECN to CE
|
||||||
|
func applyOuterECN(pkt []byte, outerECN byte, hostinfo *HostInfo, l *slog.Logger) {
|
||||||
|
if outerECN&ecnCE != ecnCE || len(pkt) < 2 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
switch pkt[0] >> 4 {
|
||||||
|
case 4:
|
||||||
|
switch pkt[1] & 0x03 {
|
||||||
|
case ecnNotECT:
|
||||||
|
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
|
||||||
|
}
|
||||||
|
case ecnCE:
|
||||||
|
// Already CE.
|
||||||
|
default:
|
||||||
|
// Rewriting the ToS byte invalidates the IPv4 header checksum, so
|
||||||
|
// patch it incrementally per RFC 1624 (HC' = ~(~HC + ~m + m')). The
|
||||||
|
// ToS is the low byte of the 16-bit word at pkt[0:2]; the header
|
||||||
|
// checksum lives at pkt[10:12]. A header too short to carry a
|
||||||
|
// checksum can't be fixed up here, so leave it for newPacket to
|
||||||
|
// reject rather than emit a mangled packet.
|
||||||
|
if len(pkt) < ipv4.HeaderLen {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
m := binary.BigEndian.Uint16(pkt[0:2])
|
||||||
|
pkt[1] = (pkt[1] &^ 0x03) | ecnCE
|
||||||
|
mNew := binary.BigEndian.Uint16(pkt[0:2])
|
||||||
|
sum := uint32(^binary.BigEndian.Uint16(pkt[10:12])) + uint32(^m) + uint32(mNew)
|
||||||
|
for sum > 0xffff {
|
||||||
|
sum = (sum >> 16) + (sum & 0xffff)
|
||||||
|
}
|
||||||
|
binary.BigEndian.PutUint16(pkt[10:12], ^uint16(sum))
|
||||||
|
}
|
||||||
|
case 6:
|
||||||
|
switch (pkt[1] >> 4) & 0x03 {
|
||||||
|
case ecnNotECT:
|
||||||
|
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
|
||||||
|
}
|
||||||
|
case ecnCE:
|
||||||
|
// Already CE.
|
||||||
|
default:
|
||||||
|
pkt[1] = (pkt[1] &^ 0x30) | (ecnCE << 4)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
|
||||||
|
// RFC 6040 normal-mode combine: fold any outer CE mark stamped by the
|
||||||
|
// underlay into the inner header before firewall + TUN write. Other
|
||||||
|
// outer codepoints are advisory only — we keep the inner unchanged.
|
||||||
|
if f.ecnEnabled.Load() {
|
||||||
|
applyOuterECN(out, meta.OuterECN, hostinfo, f.l)
|
||||||
|
}
|
||||||
|
|
||||||
err := newPacket(out, true, fwPacket)
|
err := newPacket(out, true, fwPacket)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
|
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
|
||||||
@@ -531,8 +601,10 @@ func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, p
|
|||||||
dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache)
|
dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache)
|
||||||
if dropReason != nil {
|
if dropReason != nil {
|
||||||
// NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore
|
// 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
|
// This gives us a buffer to build the reject packet in. With UDP GRO this is a single segment of a shared
|
||||||
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet, q)
|
// recvmmsg row whose capacity runs to the end of the whole row, so cap it to its own length (cap==len) to
|
||||||
|
// keep the reject builder from writing past this segment into the next, not-yet-processed coalesced segment.
|
||||||
|
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet[:len(packet):len(packet)], q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("dropping inbound packet",
|
hostinfo.logger(f.l).Debug("dropping inbound packet",
|
||||||
"fwPacket", fwPacket,
|
"fwPacket", fwPacket,
|
||||||
@@ -542,7 +614,7 @@ func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, p
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err = f.queues[q].Write(out)
|
err = f.batchers[q].Commit(out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Failed to write to tun", "error", err)
|
f.l.Error("Failed to write to tun", "error", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,28 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import "net/netip"
|
||||||
|
|
||||||
|
type RxBatcher interface {
|
||||||
|
// Reserve creates a pkt to borrow
|
||||||
|
Reserve(sz int) []byte
|
||||||
|
// Commit borrows pkt. The caller must keep pkt valid until the next Flush
|
||||||
|
Commit(pkt []byte) error
|
||||||
|
// Flush emits every queued packet in arrival order.
|
||||||
|
// Returns the first error observed; keeps draining so one bad packet doesn't hold up the rest.
|
||||||
|
// After Flush returns, borrowed payload slices may be recycled.
|
||||||
|
Flush() error
|
||||||
|
}
|
||||||
|
|
||||||
|
type TxBatcher interface {
|
||||||
|
// Reserve creates a pkt to borrow
|
||||||
|
Reserve(sz int) []byte
|
||||||
|
// Commit borrows pkt and records its destination plus the 2-bit
|
||||||
|
// IP-level ECN codepoint to set on the outer (carrier) header. The
|
||||||
|
// caller must keep pkt valid until the next Flush. Pass 0 (Not-ECT)
|
||||||
|
// to leave the outer ECN field unset.
|
||||||
|
Commit(pkt []byte, dst netip.AddrPort, outerECN byte)
|
||||||
|
// Flush emits every queued packet via the underlying batch writer in arrival order.
|
||||||
|
// Returns an errors.Join of one or more errors.
|
||||||
|
// After Flush returns, borrowed payload slices may be recycled.
|
||||||
|
Flush() error
|
||||||
|
}
|
||||||
@@ -0,0 +1,181 @@
|
|||||||
|
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 encrypted UDP socket.
|
||||||
|
const initialSlots = 64
|
||||||
|
|
||||||
|
// parsedIP is the IP-level result of parseIPPrologue. The caller layers
|
||||||
|
// L4-specific parsing (TCP / UDP) on top.
|
||||||
|
type parsedIP struct {
|
||||||
|
fk flowKey
|
||||||
|
ipHdrLen int
|
||||||
|
// pkt is the original buffer trimmed to the IP-declared total length.
|
||||||
|
// Anything below the IP layer (transport parsers) should slice into
|
||||||
|
// pkt rather than the unbounded original.
|
||||||
|
pkt []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseIPPrologue extracts the IP-level fields the coalescers care about:
|
||||||
|
// IHL/payload length, version, src/dst addresses, and the L4 protocol byte.
|
||||||
|
// Returns ok=false for malformed input, IPv4 with options or fragmentation,
|
||||||
|
// or IPv6 with extension headers (all rejected by both coalescers in
|
||||||
|
// identical ways before this refactor).
|
||||||
|
//
|
||||||
|
// On success, p.pkt is len-trimmed to the IP-declared length so callers
|
||||||
|
// don't have to repeat the trim. wantProto is the IANA protocol number to
|
||||||
|
// require (6 for TCP, 17 for UDP); ok=false for any other value.
|
||||||
|
func parseIPPrologue(pkt []byte, wantProto byte) (parsedIP, bool) {
|
||||||
|
var p parsedIP
|
||||||
|
if len(pkt) < 20 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
v := pkt[0] >> 4
|
||||||
|
switch v {
|
||||||
|
case 4:
|
||||||
|
ihl := int(pkt[0]&0x0f) * 4
|
||||||
|
if ihl != 20 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
if pkt[9] != wantProto {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
// Reject actual fragmentation (MF or non-zero frag offset).
|
||||||
|
if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
totalLen := int(binary.BigEndian.Uint16(pkt[2:4]))
|
||||||
|
if totalLen > len(pkt) || totalLen < ihl {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
p.ipHdrLen = 20
|
||||||
|
p.fk.isV6 = false
|
||||||
|
copy(p.fk.src[:4], pkt[12:16])
|
||||||
|
copy(p.fk.dst[:4], pkt[16:20])
|
||||||
|
p.pkt = pkt[:totalLen]
|
||||||
|
case 6:
|
||||||
|
if len(pkt) < 40 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
if pkt[6] != wantProto {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
payloadLen := int(binary.BigEndian.Uint16(pkt[4:6]))
|
||||||
|
if 40+payloadLen > len(pkt) {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
p.ipHdrLen = 40
|
||||||
|
p.fk.isV6 = true
|
||||||
|
copy(p.fk.src[:], pkt[8:24])
|
||||||
|
copy(p.fk.dst[:], pkt[24:40])
|
||||||
|
p.pkt = pkt[:40+payloadLen]
|
||||||
|
default:
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
return p, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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: byte 0 = version/TC[7:4], byte 1 = TC[3:0]/flow[19:16],
|
||||||
|
// bytes [2:4] = flow[15:0], [6:8] = next_hdr/hop, [8:40] = src+dst.
|
||||||
|
// Compare byte 1 fully so ECN (TC[1:0]) must match. Skip [4:6] payload_len.
|
||||||
|
if a[0] != b[0] {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if a[1] != b[1] {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[2:4], b[2:4]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[6:40], b[6:40]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
// IPv4: byte 0 = version/IHL, byte 1 = DSCP(6)|ECN(2),
|
||||||
|
// [6:10] flags/fragoff/TTL/proto, [12:20] src+dst.
|
||||||
|
// Compare byte 1 fully so ECN must match.
|
||||||
|
// Skip [2:4] total len, [4:6] id, [10:12] csum.
|
||||||
|
if a[0] != b[0] {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if a[1] != b[1] {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[6:10], b[6:10]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[12:20], b[12:20]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Arena is an injectable byte-slab that hands out non-overlapping borrowed
|
||||||
|
// 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. Pass 0 if you don't intend to call Reserve (e.g. a test that
|
||||||
|
// only feeds the coalescer pre-made []byte packets via Commit).
|
||||||
|
func NewArena(capacity int) *Arena {
|
||||||
|
return &Arena{buf: make([]byte, 0, capacity)}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reserve hands out a non-overlapping sz-byte slice from the arena. If the
|
||||||
|
// request doesn't fit the current backing, a fresh, larger backing is
|
||||||
|
// allocated; already-borrowed slices reference the old backing and remain
|
||||||
|
// valid until Reset.
|
||||||
|
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]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reserver hands out an sz-byte slice valid until its Resetter runs.
|
||||||
|
type Reserver func(sz int) []byte
|
||||||
|
|
||||||
|
// Resetter clears all reservations held by a Reserver. Only the arena's
|
||||||
|
// owner holds one; lanes inside a MultiCoalescer get nil.
|
||||||
|
type Resetter func()
|
||||||
@@ -0,0 +1,132 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
)
|
||||||
|
|
||||||
|
// MultiCoalescer fans plaintext packets out to lane-specific batchers based
|
||||||
|
// on the IP/L4 protocol of the packet, sharing a single Reserve arena
|
||||||
|
// across lanes so the caller's allocation pattern is unchanged.
|
||||||
|
//
|
||||||
|
// Lanes are processed independently: the TCP coalescer only sees TCP, the
|
||||||
|
// UDP coalescer only sees UDP, and the passthrough lane handles everything
|
||||||
|
// else. Per-flow arrival order is preserved because a single 5-tuple only
|
||||||
|
// ever lands in one lane and each lane preserves its own slot order.
|
||||||
|
//
|
||||||
|
// Cross-lane order is NOT preserved across the TCP/UDP/passthrough split.
|
||||||
|
// This is acceptable because the carrier-side recvmmsg path already
|
||||||
|
// stable-sorts by (peer, message counter) before delivering plaintext
|
||||||
|
// here, so replay-window invariants are unaffected, and apps observe
|
||||||
|
// correct per-flow ordering — which is all the IP layer guarantees anyway.
|
||||||
|
// Do not "fix" this by interleaving lane outputs at flush time; that
|
||||||
|
// negates the entire point of coalescing (each lane needs to see runs of
|
||||||
|
// adjacent same-flow packets to coalesce them).
|
||||||
|
type MultiCoalescer struct {
|
||||||
|
tcp *TCPCoalescer
|
||||||
|
udp *UDPCoalescer
|
||||||
|
pt *Passthrough
|
||||||
|
// arena is owned by the Multi: lanes get only its Reserve (nil Resetter)
|
||||||
|
// and Flush resets it exactly once after every lane has drained.
|
||||||
|
arena *Arena
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultMultiArenaCap is the recommended arena capacity for a Multi-lane
|
||||||
|
// batcher: 64 slots × 65535 bytes ≈ 4 MiB, enough to hold one recvmmsg
|
||||||
|
// burst worth of MTU-sized packets without the arena growing.
|
||||||
|
const DefaultMultiArenaCap = initialSlots * 65535
|
||||||
|
|
||||||
|
// NewMultiCoalescer builds a multi-lane batcher. tcpEnabled lets the caller
|
||||||
|
// opt out of TCP coalescing (e.g. when the queue can't do TSO); udpEnabled
|
||||||
|
// likewise gates UDP coalescing (only enable when USO was negotiated).
|
||||||
|
// Either lane disabled redirects its traffic into the passthrough lane.
|
||||||
|
// arena is the single backing slab shared across every lane; the caller
|
||||||
|
// pre-sizes it via NewArena so the hot path never allocates.
|
||||||
|
func NewMultiCoalescer(w io.Writer, l *slog.Logger, arena *Arena, tcpEnabled, udpEnabled bool) *MultiCoalescer {
|
||||||
|
m := &MultiCoalescer{
|
||||||
|
pt: NewPassthrough(w, arena.Reserve, nil),
|
||||||
|
arena: arena,
|
||||||
|
}
|
||||||
|
if tcpEnabled {
|
||||||
|
m.tcp = NewTCPCoalescer(w, l, arena.Reserve, nil)
|
||||||
|
}
|
||||||
|
if udpEnabled {
|
||||||
|
m.udp = NewUDPCoalescer(w, arena.Reserve, nil)
|
||||||
|
}
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *MultiCoalescer) Reserve(sz int) []byte {
|
||||||
|
return m.arena.Reserve(sz)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Commit dispatches pkt to the appropriate lane based on IP version + L4
|
||||||
|
// proto. Borrowed slice contract is identical to the single-lane batchers,
|
||||||
|
// pkt must remain valid until the next Flush.
|
||||||
|
//
|
||||||
|
// On the success path the IP/TCP-or-UDP parse happens here once and the
|
||||||
|
// parsed struct is handed to the lane via commitParsed so the lane doesn't
|
||||||
|
// re-walk the header.
|
||||||
|
func (m *MultiCoalescer) Commit(pkt []byte) error {
|
||||||
|
if len(pkt) < 20 {
|
||||||
|
return m.pt.Commit(pkt)
|
||||||
|
}
|
||||||
|
v := pkt[0] >> 4
|
||||||
|
var proto byte
|
||||||
|
switch v {
|
||||||
|
case 4:
|
||||||
|
proto = pkt[9]
|
||||||
|
case 6:
|
||||||
|
if len(pkt) < 40 {
|
||||||
|
return m.pt.Commit(pkt)
|
||||||
|
}
|
||||||
|
proto = pkt[6]
|
||||||
|
default:
|
||||||
|
return m.pt.Commit(pkt)
|
||||||
|
}
|
||||||
|
switch proto {
|
||||||
|
case ipProtoTCP:
|
||||||
|
if m.tcp != nil {
|
||||||
|
info, ok := parseTCPBase(pkt)
|
||||||
|
if !ok {
|
||||||
|
// Malformed/unsupported TCP shape (IP options, fragments, ...).
|
||||||
|
// Handle this via passthrough support in the TCP coalescer, to attempt to preserve flow order.
|
||||||
|
m.tcp.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return m.tcp.commitParsed(pkt, info)
|
||||||
|
}
|
||||||
|
case ipProtoUDP:
|
||||||
|
if m.udp != nil {
|
||||||
|
info, ok := parseUDP(pkt)
|
||||||
|
if !ok {
|
||||||
|
m.udp.addPassthrough(pkt) //we could also m.pt.Commit() here I guess?
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return m.udp.commitParsed(pkt, info)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return m.pt.Commit(pkt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Flush drains every lane in a fixed order, then resets the shared arena once.
|
||||||
|
// A lane error doesn't stop the remaining lanes; the joined errors are returned.
|
||||||
|
func (m *MultiCoalescer) Flush() error {
|
||||||
|
var errs []error
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
m.arena.Reset()
|
||||||
|
return errors.Join(errs...)
|
||||||
|
}
|
||||||
@@ -0,0 +1,96 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/test"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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 := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, true)
|
||||||
|
|
||||||
|
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)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(buildTCPv4(2200, tcpAck, tcpPay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(buildUDPv4(2000, 53, udpPay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(buildUDPv4(2000, 53, udpPay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(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))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestMultiCoalescerDisabledUDPFallsThrough verifies that when the UDP lane
|
||||||
|
// is disabled (e.g. kernel doesn't support USO), UDP packets still reach
|
||||||
|
// the kernel via the passthrough lane rather than being lost.
|
||||||
|
func TestMultiCoalescerDisabledUDPFallsThrough(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, false) // TSO on, USO off
|
||||||
|
|
||||||
|
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(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))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestMultiCoalescerDisabledTCPFallsThrough mirrors the TSO=off case.
|
||||||
|
func TestMultiCoalescerDisabledTCPFallsThrough(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), false, true) // TSO off, USO on
|
||||||
|
|
||||||
|
pay := make([]byte, 1200)
|
||||||
|
if err := m.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(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))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,63 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Passthrough is a RxBatcher that doesn't batch anything, it just accumulates and then sends packets.
|
||||||
|
type Passthrough struct {
|
||||||
|
out io.Writer
|
||||||
|
slots [][]byte
|
||||||
|
reserver Reserver
|
||||||
|
resetter Resetter
|
||||||
|
cursor int
|
||||||
|
}
|
||||||
|
|
||||||
|
const passthroughBaseNumSlots = 128
|
||||||
|
|
||||||
|
// DefaultPassthroughArenaCap is the recommended arena capacity for a
|
||||||
|
// standalone Passthrough batcher: 128 slots × udp.MTU ≈ 1.1 MiB.
|
||||||
|
const DefaultPassthroughArenaCap = passthroughBaseNumSlots * udp.MTU
|
||||||
|
|
||||||
|
func NewPassthrough(w io.Writer, reserver Reserver, resetter Resetter) *Passthrough {
|
||||||
|
return &Passthrough{
|
||||||
|
out: w,
|
||||||
|
slots: make([][]byte, 0, passthroughBaseNumSlots),
|
||||||
|
reserver: reserver,
|
||||||
|
resetter: resetter,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Passthrough) Reserve(sz int) []byte {
|
||||||
|
return p.reserver(sz)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Passthrough) Commit(pkt []byte) error {
|
||||||
|
p.slots = append(p.slots, pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Flush drains every queued packet and calls the configured Resetter
|
||||||
|
func (p *Passthrough) Flush() error {
|
||||||
|
firstErr := p.drain()
|
||||||
|
if p.resetter != nil {
|
||||||
|
p.resetter()
|
||||||
|
}
|
||||||
|
return firstErr
|
||||||
|
}
|
||||||
|
|
||||||
|
// drain writes out every queued packet and clears the slot list.
|
||||||
|
func (p *Passthrough) drain() error {
|
||||||
|
var firstErr error
|
||||||
|
for _, s := range p.slots {
|
||||||
|
_, err := p.out.Write(s)
|
||||||
|
if err != nil && firstErr == nil {
|
||||||
|
firstErr = err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
clear(p.slots)
|
||||||
|
p.slots = p.slots[:0]
|
||||||
|
return firstErr
|
||||||
|
}
|
||||||
@@ -0,0 +1,728 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/binary"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"net/netip"
|
||||||
|
"slices"
|
||||||
|
|
||||||
|
"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
|
||||||
|
|
||||||
|
// tcpCoalesceHdrCap is the scratch space we copy a seed's IP+TCP header
|
||||||
|
// into. IPv6 (40) + TCP with full options (60) = 100 bytes.
|
||||||
|
const tcpCoalesceHdrCap = 100
|
||||||
|
|
||||||
|
// coalesceSlot is one entry in the coalescer's ordered event queue. When
|
||||||
|
// passthrough is true the slot holds a single borrowed packet that must be
|
||||||
|
// emitted verbatim (non-TCP, non-admissible TCP, or oversize seed). When
|
||||||
|
// passthrough is false the slot is an in-progress coalesced superpacket:
|
||||||
|
// hdrBuf is a mutable copy of the seed's IP+TCP header (we patch total
|
||||||
|
// length and pseudo-header partial at flush), and payIovs are *borrowed*
|
||||||
|
// slices from the caller's plaintext buffers — no payload is ever copied.
|
||||||
|
// The caller (listenOut) must keep those buffers alive until Flush.
|
||||||
|
type coalesceSlot struct {
|
||||||
|
passthrough bool
|
||||||
|
rawPkt []byte // borrowed when passthrough
|
||||||
|
|
||||||
|
fk flowKey
|
||||||
|
hdrBuf [tcpCoalesceHdrCap]byte
|
||||||
|
hdrLen int
|
||||||
|
ipHdrLen int
|
||||||
|
isV6 bool
|
||||||
|
gsoSize int
|
||||||
|
numSeg int
|
||||||
|
totalPay int
|
||||||
|
nextSeq uint32
|
||||||
|
// psh closes the chain: set when the last-accepted segment had PSH or
|
||||||
|
// was sub-gsoSize. No further appends after that.
|
||||||
|
psh bool
|
||||||
|
payIovs [][]byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// TCPCoalescer accumulates adjacent in-flow TCP data segments across
|
||||||
|
// multiple concurrent flows and emits each flow's run as a single TSO
|
||||||
|
// superpacket via tio.GSOWriter. All output — coalesced or not — is
|
||||||
|
// deferred until Flush so arrival order is preserved on the wire. Owns
|
||||||
|
// no locks; one coalescer per TUN write queue.
|
||||||
|
type TCPCoalescer struct {
|
||||||
|
plainW io.Writer
|
||||||
|
gsoW tio.GSOWriter // nil when the queue doesn't support TSO
|
||||||
|
|
||||||
|
// slots is the ordered event queue. Flush walks it once and emits each
|
||||||
|
// entry as either a WriteGSO (coalesced) or a plainW.Write (passthrough).
|
||||||
|
slots []*coalesceSlot
|
||||||
|
// openSlots maps a flow key to its most recent non-sealed slot, so new
|
||||||
|
// segments can extend an in-progress superpacket in O(1). Slots are
|
||||||
|
// removed from this map when they close (PSH or short-last-segment),
|
||||||
|
// when a non-admissible packet for that flow arrives, or in Flush.
|
||||||
|
openSlots map[flowKey]*coalesceSlot
|
||||||
|
// lastSlot caches the most recently touched open slot. Steady-state
|
||||||
|
// bulk traffic is dominated by a single flow, so comparing the
|
||||||
|
// incoming key against the cached slot's own fk lets the hot path
|
||||||
|
// skip the map lookup (and the aeshash of a 38-byte key) entirely.
|
||||||
|
// Kept in lockstep with openSlots: nil whenever the slot it pointed
|
||||||
|
// at is removed/sealed.
|
||||||
|
lastSlot *coalesceSlot
|
||||||
|
pool []*coalesceSlot // free list for reuse
|
||||||
|
reserver Reserver
|
||||||
|
resetter Resetter
|
||||||
|
l *slog.Logger
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewTCPCoalescer(w io.Writer, l *slog.Logger, reserver Reserver, resetter Resetter) *TCPCoalescer {
|
||||||
|
c := &TCPCoalescer{
|
||||||
|
plainW: w,
|
||||||
|
slots: make([]*coalesceSlot, 0, initialSlots),
|
||||||
|
openSlots: make(map[flowKey]*coalesceSlot, initialSlots),
|
||||||
|
pool: make([]*coalesceSlot, 0, initialSlots),
|
||||||
|
reserver: reserver,
|
||||||
|
resetter: resetter,
|
||||||
|
l: l,
|
||||||
|
}
|
||||||
|
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP); ok {
|
||||||
|
c.gsoW = gw
|
||||||
|
}
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// parsedTCP holds the fields extracted from a single parse so later steps
|
||||||
|
// (admission, slot lookup, canAppend) don't re-walk the header.
|
||||||
|
type parsedTCP struct {
|
||||||
|
fk flowKey
|
||||||
|
ipHdrLen int
|
||||||
|
tcpHdrLen int
|
||||||
|
hdrLen int
|
||||||
|
payLen int
|
||||||
|
seq uint32
|
||||||
|
flags byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseTCPBase extracts the flow key and IP/TCP offsets for any TCP packet,
|
||||||
|
// regardless of whether it's admissible for coalescing. Returns ok=false
|
||||||
|
// for non-TCP or malformed input.
|
||||||
|
// Accepts IPv4 (no options or fragmentation) and IPv6 (no extension headers).
|
||||||
|
func parseTCPBase(pkt []byte) (parsedTCP, bool) {
|
||||||
|
var p parsedTCP
|
||||||
|
ip, ok := parseIPPrologue(pkt, ipProtoTCP)
|
||||||
|
if !ok {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
pkt = ip.pkt
|
||||||
|
p.fk = ip.fk
|
||||||
|
p.ipHdrLen = ip.ipHdrLen
|
||||||
|
|
||||||
|
if len(pkt) < p.ipHdrLen+20 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
tcpOff := int(pkt[p.ipHdrLen+12]>>4) * 4
|
||||||
|
if tcpOff < 20 || tcpOff > 60 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
if len(pkt) < p.ipHdrLen+tcpOff {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
p.tcpHdrLen = tcpOff
|
||||||
|
p.hdrLen = p.ipHdrLen + tcpOff
|
||||||
|
p.payLen = len(pkt) - p.hdrLen
|
||||||
|
p.seq = binary.BigEndian.Uint32(pkt[p.ipHdrLen+4 : p.ipHdrLen+8])
|
||||||
|
p.flags = pkt[p.ipHdrLen+13]
|
||||||
|
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2])
|
||||||
|
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
|
||||||
|
return p, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// TCP flag bits (byte 13 of the TCP header). Only the bits actually consulted
|
||||||
|
// by the coalescer are named; FIN/SYN/RST/URG/CWR are rejected via the
|
||||||
|
// negative mask in coalesceable, not by name.
|
||||||
|
const (
|
||||||
|
tcpFlagPsh = 0x08
|
||||||
|
tcpFlagAck = 0x10
|
||||||
|
tcpFlagEce = 0x40
|
||||||
|
)
|
||||||
|
|
||||||
|
// coalesceable reports whether a parsed TCP segment is eligible for
|
||||||
|
// coalescing. Accepts ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE with a
|
||||||
|
// non-empty payload. CWR is excluded because it marks a one-shot
|
||||||
|
// congestion-window-reduced transition the receiver must observe at a
|
||||||
|
// segment boundary.
|
||||||
|
func (p parsedTCP) coalesceable() bool {
|
||||||
|
if p.flags&tcpFlagAck == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if p.flags&^(tcpFlagAck|tcpFlagPsh|tcpFlagEce) != 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return p.payLen > 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TCPCoalescer) Reserve(sz int) []byte {
|
||||||
|
return c.reserver(sz)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
||||||
|
func (c *TCPCoalescer) Commit(pkt []byte) error {
|
||||||
|
if c.gsoW == nil {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
info, ok := parseTCPBase(pkt)
|
||||||
|
if !ok {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return c.commitParsed(pkt, info)
|
||||||
|
}
|
||||||
|
|
||||||
|
// commitParsed is the post-parse half of Commit. The caller must have
|
||||||
|
// already verified parseTCPBase succeeded (info is a valid TCP parse).
|
||||||
|
// Used by MultiCoalescer.Commit to avoid re-walking the IP/TCP header
|
||||||
|
// after the dispatcher has already done so.
|
||||||
|
func (c *TCPCoalescer) commitParsed(pkt []byte, info parsedTCP) error {
|
||||||
|
if c.gsoW == nil {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if !info.coalesceable() {
|
||||||
|
// TCP but not admissible (SYN/FIN/RST/URG/CWR or zero-payload).
|
||||||
|
// Seal this flow's open slot so later in-flow packets don't extend
|
||||||
|
// it and accidentally reorder past this passthrough.
|
||||||
|
if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||||
|
c.lastSlot = nil
|
||||||
|
}
|
||||||
|
delete(c.openSlots, info.fk)
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Single-flow fast path: with only one open flow the cache hits every
|
||||||
|
// packet, and len(openSlots)==1 lets us skip the 38-byte fk compare
|
||||||
|
// when there are multiple flows in flight (where the hit rate would
|
||||||
|
// be ~0 and the compare is pure overhead).
|
||||||
|
var open *coalesceSlot
|
||||||
|
if last := c.lastSlot; last != nil && len(c.openSlots) == 1 && last.fk == info.fk {
|
||||||
|
open = last
|
||||||
|
} else {
|
||||||
|
open = c.openSlots[info.fk]
|
||||||
|
}
|
||||||
|
if open != nil {
|
||||||
|
if c.canAppend(open, pkt, info) {
|
||||||
|
c.appendPayload(open, pkt, info)
|
||||||
|
if open.psh {
|
||||||
|
delete(c.openSlots, info.fk)
|
||||||
|
c.lastSlot = nil
|
||||||
|
} else {
|
||||||
|
c.lastSlot = open
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
// Can't extend — seal it and fall through to seed a fresh slot.
|
||||||
|
delete(c.openSlots, info.fk)
|
||||||
|
if c.lastSlot == open {
|
||||||
|
c.lastSlot = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
c.seed(pkt, info)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Flush emits every queued event in (per-flow) seq order.
|
||||||
|
func (c *TCPCoalescer) Flush() error {
|
||||||
|
first := c.drain()
|
||||||
|
if c.resetter != nil {
|
||||||
|
c.resetter()
|
||||||
|
}
|
||||||
|
return first
|
||||||
|
}
|
||||||
|
|
||||||
|
// drain emits every queued slot (reordering/merging coalesced runs first)
|
||||||
|
// and clears the slot state.
|
||||||
|
func (c *TCPCoalescer) drain() error {
|
||||||
|
c.reorderForFlush()
|
||||||
|
var first error
|
||||||
|
for _, s := range c.slots {
|
||||||
|
var err error
|
||||||
|
if s.passthrough {
|
||||||
|
_, err = c.plainW.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) addPassthrough(pkt []byte) {
|
||||||
|
s := c.take()
|
||||||
|
s.passthrough = true
|
||||||
|
s.rawPkt = pkt
|
||||||
|
c.slots = append(c.slots, s)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TCPCoalescer) seed(pkt []byte, info parsedTCP) {
|
||||||
|
if info.hdrLen > tcpCoalesceHdrCap || info.hdrLen+info.payLen > tcpCoalesceBufSize {
|
||||||
|
// Pathological shape — can't fit our scratch, emit as-is.
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s := c.take()
|
||||||
|
s.passthrough = false
|
||||||
|
s.rawPkt = nil
|
||||||
|
copy(s.hdrBuf[:], pkt[:info.hdrLen])
|
||||||
|
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.psh = info.flags&tcpFlagPsh != 0
|
||||||
|
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||||
|
c.slots = append(c.slots, s)
|
||||||
|
if !s.psh {
|
||||||
|
c.openSlots[info.fk] = s
|
||||||
|
c.lastSlot = s
|
||||||
|
} else if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||||
|
// PSH-on-seed seals the slot immediately. Any prior cached open
|
||||||
|
// slot for this flow has just been sealed-and-replaced by this
|
||||||
|
// passthrough-shaped seed, so drop the cache too.
|
||||||
|
c.lastSlot = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// canAppend reports whether info's packet extends the slot's seed: same
|
||||||
|
// header shape and stable contents, adjacent seq, not oversized, chain not closed.
|
||||||
|
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info parsedTCP) bool {
|
||||||
|
if s.psh {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if info.hdrLen != s.hdrLen {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
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.hdrBuf[s.ipHdrLen+13]
|
||||||
|
if (seedFlags^info.flags)&tcpFlagEce != 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !headersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info parsedTCP) {
|
||||||
|
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||||
|
s.numSeg++
|
||||||
|
s.totalPay += info.payLen
|
||||||
|
s.nextSeq = info.seq + uint32(info.payLen)
|
||||||
|
if info.flags&tcpFlagPsh != 0 {
|
||||||
|
// Propagate PSH into the seed header so kernel TSO sets it on the
|
||||||
|
// last segment. Without this the sender's push signal is dropped.
|
||||||
|
s.hdrBuf[s.ipHdrLen+13] |= tcpFlagPsh
|
||||||
|
}
|
||||||
|
if info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0 {
|
||||||
|
s.psh = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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) {
|
||||||
|
s.passthrough = false
|
||||||
|
s.rawPkt = nil
|
||||||
|
clear(s.payIovs)
|
||||||
|
s.payIovs = s.payIovs[:0]
|
||||||
|
s.numSeg = 0
|
||||||
|
s.totalPay = 0
|
||||||
|
s.psh = false
|
||||||
|
c.pool = append(c.pool, s)
|
||||||
|
}
|
||||||
|
|
||||||
|
// flushSlot patches the header and calls WriteGSO. Does not remove the slot from c.slots.
|
||||||
|
func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
|
||||||
|
total := s.hdrLen + s.totalPay
|
||||||
|
l4Len := total - s.ipHdrLen
|
||||||
|
hdr := s.hdrBuf[: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.gsoW.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
|
||||||
|
}
|
||||||
|
|
||||||
|
// reorderForFlush neutralizes wire-side reorder that the rxOrder buffer
|
||||||
|
// couldn't catch (anything crossing a recvmmsg batch boundary). Without
|
||||||
|
// this pass a small wire reorder — counter 250 arriving in batch K when
|
||||||
|
// 200..249 are coming in batch K+1 — would seed an out-of-seq slot first
|
||||||
|
// and emit it ahead of the lower-seq slot, manifesting at the inner TCP
|
||||||
|
// receiver as a much larger reorder than the wire actually had.
|
||||||
|
//
|
||||||
|
// Two phases:
|
||||||
|
// 1. Sort each passthrough-bounded segment of c.slots by (flow, seq).
|
||||||
|
// Cross-flow ordering inside a segment isn't preserved (it never was
|
||||||
|
// and doesn't matter for any single flow's TCP correctness).
|
||||||
|
// 2. Sweep once and merge adjacent same-flow slots whose ranges are now
|
||||||
|
// contiguous AND whose tail is gsoSize-aligned. The tail constraint
|
||||||
|
// matters because the kernel TSO splitter chops at gsoSize from the
|
||||||
|
// start of the merged payload — a short segment in the middle would
|
||||||
|
// desynchronize every later segment.
|
||||||
|
//
|
||||||
|
// Passthrough slots act as barriers: the merge check skips them on either
|
||||||
|
// side, so a SYN/FIN/RST/CWR is never reordered relative to its flow's
|
||||||
|
// data.
|
||||||
|
func (c *TCPCoalescer) reorderForFlush() {
|
||||||
|
if len(c.slots) <= 1 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
runStart := 0
|
||||||
|
for i := 0; i <= len(c.slots); i++ {
|
||||||
|
if i < len(c.slots) && !c.slots[i].passthrough {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
c.sortRun(c.slots[runStart:i])
|
||||||
|
runStart = i + 1
|
||||||
|
}
|
||||||
|
out := c.slots[:0]
|
||||||
|
logged := false
|
||||||
|
for _, s := range c.slots {
|
||||||
|
if n := len(out); n > 0 {
|
||||||
|
prev := out[n-1]
|
||||||
|
if !prev.passthrough && !s.passthrough && prev.fk == s.fk {
|
||||||
|
// Same-flow neighbors after sort. If they aren't seq-
|
||||||
|
// contiguous it's a real gap — packets the wire reordered
|
||||||
|
// across batches, or actual loss before nebula. Log it so
|
||||||
|
// the operator can quantify how often it happens; the data
|
||||||
|
// itself still emits in seq order, kernel TCP handles the
|
||||||
|
// gap via its OOO queue.
|
||||||
|
if c.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
if prev.nextSeq != slotSeedSeq(s) {
|
||||||
|
logged = true
|
||||||
|
gap := int64(slotSeedSeq(s)) - int64(prev.nextSeq)
|
||||||
|
c.l.Debug("tcp coalesce: cross-slot seq gap",
|
||||||
|
"src", flowKeyAddr(s.fk, false),
|
||||||
|
"dst", flowKeyAddr(s.fk, true),
|
||||||
|
"sport", s.fk.sport,
|
||||||
|
"dport", s.fk.dport,
|
||||||
|
"prev_seed_seq", slotSeedSeq(prev),
|
||||||
|
"prev_next_seq", prev.nextSeq,
|
||||||
|
"this_seed_seq", slotSeedSeq(s),
|
||||||
|
"gap_bytes", gap,
|
||||||
|
"prev_seg_count", prev.numSeg,
|
||||||
|
"prev_total_pay", prev.totalPay,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if canMergeSlots(prev, s) {
|
||||||
|
mergeSlots(prev, s)
|
||||||
|
c.release(s)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
out = append(out, s)
|
||||||
|
}
|
||||||
|
if logged {
|
||||||
|
c.l.Warn("==== end of batch ====")
|
||||||
|
}
|
||||||
|
c.slots = out
|
||||||
|
}
|
||||||
|
|
||||||
|
// flowKeyAddr returns the src or dst address from fk as a netip.Addr for
|
||||||
|
// logging. Only used on the cold gap-log path so the netip allocation
|
||||||
|
// doesn't matter.
|
||||||
|
func flowKeyAddr(fk flowKey, dst bool) netip.Addr {
|
||||||
|
src := fk.src
|
||||||
|
if dst {
|
||||||
|
src = fk.dst
|
||||||
|
}
|
||||||
|
if fk.isV6 {
|
||||||
|
return netip.AddrFrom16(src)
|
||||||
|
}
|
||||||
|
var v4 [4]byte
|
||||||
|
copy(v4[:], src[:4])
|
||||||
|
return netip.AddrFrom4(v4)
|
||||||
|
}
|
||||||
|
|
||||||
|
// sortRun stable-sorts run by (flowKey, seedSeq) so each flow's slots
|
||||||
|
// cluster together in seq order, ready for the merge sweep. Stable so
|
||||||
|
// equal-key slots keep their original relative position (defensive — a
|
||||||
|
// duplicate seedSeq would already mean something's wrong upstream).
|
||||||
|
func (c *TCPCoalescer) sortRun(run []*coalesceSlot) {
|
||||||
|
if len(run) <= 1 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
// slices.SortStableFunc with a free, non-capturing comparator avoids the
|
||||||
|
// reflection + closure-escape allocations that sort.SliceStable forces.
|
||||||
|
slices.SortStableFunc(run, compareCoalesceSlots)
|
||||||
|
}
|
||||||
|
|
||||||
|
func compareCoalesceSlots(a, b *coalesceSlot) int {
|
||||||
|
if cmp := flowKeyCompare(a.fk, b.fk); cmp != 0 {
|
||||||
|
return cmp
|
||||||
|
}
|
||||||
|
aSeq, bSeq := slotSeedSeq(a), slotSeedSeq(b)
|
||||||
|
if aSeq == bSeq {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
if tcpSeqLess(aSeq, bSeq) {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
|
||||||
|
// slotSeedSeq returns the TCP seq of the slot's seed (first segment).
|
||||||
|
// nextSeq tracks the seq just past the last appended byte; subtracting
|
||||||
|
// totalPay walks back to the seed. uint32 wraparound is the right TCP
|
||||||
|
// arithmetic so no special-casing is needed.
|
||||||
|
func slotSeedSeq(s *coalesceSlot) uint32 {
|
||||||
|
return s.nextSeq - uint32(s.totalPay)
|
||||||
|
}
|
||||||
|
|
||||||
|
// tcpSeqLess reports whether a precedes b in TCP serial-number arithmetic
|
||||||
|
// (RFC 1323 §2.3). The signed int32 cast turns the modular subtraction
|
||||||
|
// into the right comparison even across the 2^32 wrap.
|
||||||
|
func tcpSeqLess(a, b uint32) bool {
|
||||||
|
return int32(a-b) < 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// flowKeyCompare orders flowKeys deterministically. The exact ordering
|
||||||
|
// is irrelevant — only that same-flow slots cluster together so the
|
||||||
|
// post-sort sweep can merge contiguous pairs.
|
||||||
|
func flowKeyCompare(a, b flowKey) int {
|
||||||
|
// Cheap scalar fields first so most non-matching keys short-circuit
|
||||||
|
// without ever calling bytes.Compare. sport is the ephemeral port on
|
||||||
|
// egress flows and discriminates fastest. For matching keys (same
|
||||||
|
// flow), array equality on src/dst inlines to word-sized compares,
|
||||||
|
// so we only pay bytes.Compare when the arrays actually differ.
|
||||||
|
if a.sport != b.sport {
|
||||||
|
if a.sport < b.sport {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
if a.dport != b.dport {
|
||||||
|
if a.dport < b.dport {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
if a.dst != b.dst {
|
||||||
|
return bytes.Compare(a.dst[:], b.dst[:])
|
||||||
|
}
|
||||||
|
if a.src != b.src {
|
||||||
|
return bytes.Compare(a.src[:], b.src[:])
|
||||||
|
}
|
||||||
|
if a.isV6 != b.isV6 {
|
||||||
|
if !a.isV6 {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// canMergeSlots reports whether s can fold into prev as one merged TSO
|
||||||
|
// superpacket. Same flow, contiguous TCP byte range, equal gsoSize, and
|
||||||
|
// fits within the kernel TSO limits. The tail-of-prev check rejects any
|
||||||
|
// merge whose first slot ended on a sub-gsoSize segment — kernel TSO
|
||||||
|
// would split the merged skb at gsoSize boundaries from the start, so a
|
||||||
|
// short segment in the middle would corrupt every later segment. PSH and
|
||||||
|
// ECE state must agree across both slots: PSH is a semantic delimiter
|
||||||
|
// (preserving the sender's push boundary) and ECE state must be uniform
|
||||||
|
// across a window (the same rule canAppend enforces for in-flow appends).
|
||||||
|
// The IP-level ECN codepoint must also match: this check calls headersMatch
|
||||||
|
// → ipHeadersMatch, which compares the full DSCP/ECN byte, so two slots with
|
||||||
|
// differing ECN marks stay separate superpackets, each keeping its own mark.
|
||||||
|
//
|
||||||
|
// Note: a slot sealed by reorder (canAppend returned false on seq
|
||||||
|
// mismatch) keeps psh=false, so this restriction does not block the
|
||||||
|
// reorder-fix merge — only legitimate PSH-set seals.
|
||||||
|
func canMergeSlots(prev, s *coalesceSlot) bool {
|
||||||
|
if prev.psh {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if prev.fk != s.fk {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if prev.gsoSize != s.gsoSize {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if prev.nextSeq != slotSeedSeq(s) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if prev.numSeg+s.numSeg > tcpCoalesceMaxSegs {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if prev.hdrLen+prev.totalPay+s.totalPay > tcpCoalesceBufSize {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if len(prev.payIovs[len(prev.payIovs)-1]) != prev.gsoSize {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
prevFlags := prev.hdrBuf[prev.ipHdrLen+13]
|
||||||
|
sFlags := s.hdrBuf[s.ipHdrLen+13]
|
||||||
|
if (prevFlags^sFlags)&tcpFlagEce != 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !headersMatch(prev.hdrBuf[:prev.hdrLen], s.hdrBuf[:s.hdrLen], prev.isV6, prev.ipHdrLen) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// mergeSlots folds src into dst in place: payIovs concatenated, counters
|
||||||
|
// and totals updated, PSH OR'd into the seed header so the push signal is
|
||||||
|
// not lost. The seed header's seq, gsoSize, and fk are unchanged. Caller
|
||||||
|
// is responsible for releasing src (it's no longer in c.slots after this call).
|
||||||
|
func mergeSlots(dst, src *coalesceSlot) {
|
||||||
|
dst.payIovs = append(dst.payIovs, src.payIovs...)
|
||||||
|
dst.numSeg += src.numSeg
|
||||||
|
dst.totalPay += src.totalPay
|
||||||
|
dst.nextSeq = src.nextSeq
|
||||||
|
if src.psh {
|
||||||
|
dst.psh = true
|
||||||
|
dst.hdrBuf[dst.ipHdrLen+13] |= tcpFlagPsh
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must
|
||||||
|
// already have its checksum field zeroed) and returns the folded/inverted
|
||||||
|
// 16-bit value to store.
|
||||||
|
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 — the kernel will add the payload sum and invert.
|
||||||
|
func foldOnceNoInvert(sum uint32) uint16 {
|
||||||
|
for sum>>16 != 0 {
|
||||||
|
sum = (sum & 0xffff) + (sum >> 16)
|
||||||
|
}
|
||||||
|
return uint16(sum)
|
||||||
|
}
|
||||||
@@ -0,0 +1,241 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"runtime"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"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
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildICMPv4 returns a minimal non-TCP packet that takes the passthrough
|
||||||
|
// 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()
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewTCPCoalescer(nopTunWriter{}, test.NewLogger(), arena.Reserve, arena.Reset)
|
||||||
|
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))
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkCommitPassthrough exercises the non-TCP branch: parseTCPBase
|
||||||
|
// bails early and addPassthrough is the only work.
|
||||||
|
func BenchmarkCommitPassthrough(b *testing.B) {
|
||||||
|
pkt := buildICMPv4()
|
||||||
|
pkts := make([][]byte, 64)
|
||||||
|
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 + passthrough. 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. The dispatcher does
|
||||||
|
// the IP/L4 parse once and passes the parsed struct to the lane, so this
|
||||||
|
// is the bench that shows the savings of skipping the lane's re-parse.
|
||||||
|
func runMultiCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||||
|
b.Helper()
|
||||||
|
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger(), NewArena(0), true, true)
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.SetBytes(int64(len(pkts[0])))
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
pkt := pkts[i%len(pkts)]
|
||||||
|
if err := m.Commit(pkt); 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))
|
||||||
|
}
|
||||||
|
|
||||||
|
// flowKeyPair is one comparison input for the flowKeyCompare bench.
|
||||||
|
type flowKeyPair struct{ a, b flowKey }
|
||||||
|
|
||||||
|
// makeFlowKey builds an IPv4 flowKey from compact inputs.
|
||||||
|
func makeFlowKey(srcLow, dstLow uint32, sport, dport uint16) flowKey {
|
||||||
|
var fk flowKey
|
||||||
|
binary.BigEndian.PutUint32(fk.src[12:16], srcLow)
|
||||||
|
binary.BigEndian.PutUint32(fk.dst[12:16], dstLow)
|
||||||
|
fk.sport = sport
|
||||||
|
fk.dport = dport
|
||||||
|
return fk
|
||||||
|
}
|
||||||
|
|
||||||
|
// flowKeyCases are the workload mixes flowKeyCompare sees in practice.
|
||||||
|
// - sameFlow: equal keys; tests the equal-path cost (sort runs hit this
|
||||||
|
// repeatedly when many segments share a flow).
|
||||||
|
// - sportDiffers: same src/dst/dport, different sport — the typical
|
||||||
|
// "sibling flows from one host to one server" pattern.
|
||||||
|
// - dstDiffers: same src/sport/dport, different dst — outbound to many
|
||||||
|
// servers from a fixed local port.
|
||||||
|
// - allDiffer: every field differs; worst case for short-circuiting.
|
||||||
|
func flowKeyCases() map[string][]flowKeyPair {
|
||||||
|
const n = 64
|
||||||
|
cases := map[string][]flowKeyPair{
|
||||||
|
"sameFlow": make([]flowKeyPair, n),
|
||||||
|
"sportDiffers": make([]flowKeyPair, n),
|
||||||
|
"dstDiffers": make([]flowKeyPair, n),
|
||||||
|
"allDiffer": make([]flowKeyPair, n),
|
||||||
|
}
|
||||||
|
for i := range n {
|
||||||
|
base := makeFlowKey(0x0a000001, 0x0a000002, 40000, 443)
|
||||||
|
cases["sameFlow"][i] = flowKeyPair{a: base, b: base}
|
||||||
|
cases["sportDiffers"][i] = flowKeyPair{
|
||||||
|
a: base,
|
||||||
|
b: makeFlowKey(0x0a000001, 0x0a000002, uint16(40001+i), 443),
|
||||||
|
}
|
||||||
|
cases["dstDiffers"][i] = flowKeyPair{
|
||||||
|
a: base,
|
||||||
|
b: makeFlowKey(0x0a000001, uint32(0x0a000002+i+1), 40000, 443),
|
||||||
|
}
|
||||||
|
cases["allDiffer"][i] = flowKeyPair{
|
||||||
|
a: makeFlowKey(uint32(0x0a000001+i), uint32(0x0a000002+i), uint16(40000+i), uint16(80+i)),
|
||||||
|
b: makeFlowKey(uint32(0x0b000001+i), uint32(0x0b000002+i), uint16(50000+i), uint16(443+i)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return cases
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkFlowKeyCompare measures flowKeyCompare across the workloads
|
||||||
|
// the sort step actually sees. Use this to compare reorderings.
|
||||||
|
func BenchmarkFlowKeyCompare(b *testing.B) {
|
||||||
|
for name, pairs := range flowKeyCases() {
|
||||||
|
b.Run(name, func(b *testing.B) {
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
var sink int
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
p := pairs[i&(len(pairs)-1)]
|
||||||
|
sink += flowKeyCompare(p.a, p.b)
|
||||||
|
}
|
||||||
|
runtime.KeepAlive(sink)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,60 @@
|
|||||||
|
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, outerECNs []byte) 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
|
||||||
|
ecns []byte
|
||||||
|
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),
|
||||||
|
ecns: make([]byte, 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, outerECN byte) {
|
||||||
|
b.bufs = append(b.bufs, pkt)
|
||||||
|
b.dsts = append(b.dsts, dst)
|
||||||
|
b.ecns = append(b.ecns, outerECN)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *SendBatch) Flush() error {
|
||||||
|
var err error
|
||||||
|
if len(b.bufs) > 0 {
|
||||||
|
err = b.out.WriteBatch(b.bufs, b.dsts, b.ecns)
|
||||||
|
}
|
||||||
|
clear(b.bufs)
|
||||||
|
b.bufs = b.bufs[:0]
|
||||||
|
b.dsts = b.dsts[:0]
|
||||||
|
b.ecns = b.ecns[:0]
|
||||||
|
b.arena.Reset()
|
||||||
|
return err
|
||||||
|
}
|
||||||
@@ -0,0 +1,124 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
type fakeBatchWriter struct {
|
||||||
|
bufs [][]byte
|
||||||
|
addrs []netip.AddrPort
|
||||||
|
ecns []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, ecns []byte) error {
|
||||||
|
// Snapshot — SendBatch.Flush nils its slot pointers right after WriteBatch
|
||||||
|
// returns, so tests must capture data before that happens.
|
||||||
|
w.bufs = make([][]byte, len(bufs))
|
||||||
|
for i, b := range bufs {
|
||||||
|
cp := make([]byte, len(b))
|
||||||
|
copy(cp, b)
|
||||||
|
w.bufs[i] = cp
|
||||||
|
}
|
||||||
|
w.addrs = append(w.addrs[:0], addrs...)
|
||||||
|
w.ecns = append(w.ecns[:0], ecns...)
|
||||||
|
return 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, 0)
|
||||||
|
}
|
||||||
|
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, 0)
|
||||||
|
}
|
||||||
|
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, 0)
|
||||||
|
|
||||||
|
s2 := b.Reserve(8) // exceeds remaining cap, triggers grow
|
||||||
|
pkt2 := append(s2[:0], 0xA, 0xB, 0xC, 0xD, 0xE)
|
||||||
|
b.Commit(pkt2, ap, 0)
|
||||||
|
|
||||||
|
// pkt1 must still be intact even though backing reallocated.
|
||||||
|
if pkt1[0] != 0x11 || pkt1[3] != 0x44 {
|
||||||
|
t.Fatalf("first packet corrupted by grow: %x", pkt1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := b.Flush(); err != nil {
|
||||||
|
t.Fatalf("Flush: %v", err)
|
||||||
|
}
|
||||||
|
if len(fw.bufs) != 2 {
|
||||||
|
t.Fatalf("got %d bufs want 2", len(fw.bufs))
|
||||||
|
}
|
||||||
|
if fw.bufs[0][0] != 0x11 || fw.bufs[0][3] != 0x44 {
|
||||||
|
t.Errorf("first packet on the wire: %x", fw.bufs[0])
|
||||||
|
}
|
||||||
|
if fw.bufs[1][0] != 0xA || fw.bufs[1][4] != 0xE {
|
||||||
|
t.Errorf("second packet on the wire: %x", fw.bufs[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,355 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"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
|
||||||
|
|
||||||
|
// udpCoalesceHdrCap is the scratch space we copy a seed's IP+UDP header
|
||||||
|
// into. IPv6 (40) + UDP (8) = 48; round up for safety.
|
||||||
|
const udpCoalesceHdrCap = 64
|
||||||
|
|
||||||
|
// udpSlot is one entry in the UDPCoalescer's ordered event queue.
|
||||||
|
type udpSlot struct {
|
||||||
|
passthrough bool
|
||||||
|
rawPkt []byte // borrowed when passthrough
|
||||||
|
|
||||||
|
fk flowKey
|
||||||
|
hdrBuf [udpCoalesceHdrCap]byte
|
||||||
|
hdrLen int
|
||||||
|
ipHdrLen int
|
||||||
|
isV6 bool
|
||||||
|
gsoSize int // per-segment UDP payload length
|
||||||
|
numSeg int
|
||||||
|
totalPay int
|
||||||
|
// sealed closes the chain: set when a sub-gsoSize segment is appended
|
||||||
|
// (kernel UDP-GSO requires every segment but the last to be exactly
|
||||||
|
// gsoSize) or when limits are hit. No further appends after.
|
||||||
|
sealed bool
|
||||||
|
payIovs [][]byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// UDPCoalescer accumulates adjacent in-flow UDP datagrams across multiple
|
||||||
|
// concurrent flows and emits each flow's run as a single GSO_UDP_L4
|
||||||
|
// superpacket via tio.GSOWriter. Falls back to per-packet writes when the
|
||||||
|
// underlying writer doesn't support USO.
|
||||||
|
//
|
||||||
|
// All output — coalesced or not — is deferred until Flush so per-flow
|
||||||
|
// arrival order is preserved on the wire. Cross-flow order is NOT preserved
|
||||||
|
// across the TCP/UDP/passthrough split when this coalescer runs alongside
|
||||||
|
// others — see multi_coalesce.go. Per-flow order is preserved because a
|
||||||
|
// single 5-tuple only ever lands in one lane and each lane preserves its
|
||||||
|
// own slot order.
|
||||||
|
//
|
||||||
|
// Owns no locks; one coalescer per TUN write queue.
|
||||||
|
type UDPCoalescer struct {
|
||||||
|
plainW io.Writer
|
||||||
|
gsoW tio.GSOWriter // nil when the queue can't accept GSO_UDP_L4
|
||||||
|
|
||||||
|
slots []*udpSlot
|
||||||
|
openSlots map[flowKey]*udpSlot
|
||||||
|
pool []*udpSlot
|
||||||
|
reserver Reserver
|
||||||
|
resetter Resetter
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewUDPCoalescer wraps w. The caller is responsible for only constructing
|
||||||
|
// this when the underlying Queue's Capabilities advertise USO; otherwise
|
||||||
|
// the kernel may reject GSO_UDP_L4 writes. If w does not implement
|
||||||
|
// tio.GSOWriter at all (single-packet Queue), the coalescer degrades to
|
||||||
|
// plain Writes — same defensive shape as the TCP coalescer.
|
||||||
|
func NewUDPCoalescer(w io.Writer, reserver Reserver, resetter Resetter) *UDPCoalescer {
|
||||||
|
c := &UDPCoalescer{
|
||||||
|
plainW: w,
|
||||||
|
slots: make([]*udpSlot, 0, initialSlots),
|
||||||
|
openSlots: make(map[flowKey]*udpSlot, initialSlots),
|
||||||
|
pool: make([]*udpSlot, 0, initialSlots),
|
||||||
|
reserver: reserver,
|
||||||
|
resetter: resetter,
|
||||||
|
}
|
||||||
|
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP); ok {
|
||||||
|
c.gsoW = gw
|
||||||
|
}
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// parsedUDP holds the fields extracted from a single parse so later steps
|
||||||
|
// (admission, slot lookup, canAppend) don't re-walk the header.
|
||||||
|
type parsedUDP struct {
|
||||||
|
fk flowKey
|
||||||
|
ipHdrLen int
|
||||||
|
hdrLen int // ipHdrLen + 8
|
||||||
|
payLen int
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseUDP extracts the flow key and IP/UDP offsets for a UDP packet.
|
||||||
|
// Returns ok=false for non-UDP, malformed, or unsupported header shapes
|
||||||
|
// (IPv4 with options/fragmentation, IPv6 with extension headers).
|
||||||
|
func parseUDP(pkt []byte) (parsedUDP, bool) {
|
||||||
|
var p parsedUDP
|
||||||
|
ip, ok := parseIPPrologue(pkt, ipProtoUDP)
|
||||||
|
if !ok {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
pkt = ip.pkt
|
||||||
|
p.fk = ip.fk
|
||||||
|
p.ipHdrLen = ip.ipHdrLen
|
||||||
|
|
||||||
|
if len(pkt) < p.ipHdrLen+8 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
p.hdrLen = p.ipHdrLen + 8
|
||||||
|
// UDP `length` field: must equal IP-derived length-of-UDP-header-plus-payload.
|
||||||
|
udpLen := int(binary.BigEndian.Uint16(pkt[p.ipHdrLen+4 : p.ipHdrLen+6]))
|
||||||
|
if udpLen < 8 || udpLen > len(pkt)-p.ipHdrLen {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
p.payLen = udpLen - 8
|
||||||
|
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2])
|
||||||
|
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
|
||||||
|
return p, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCoalescer) Reserve(sz int) []byte {
|
||||||
|
return c.reserver(sz)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
||||||
|
func (c *UDPCoalescer) Commit(pkt []byte) error {
|
||||||
|
if c.gsoW == nil {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
info, ok := parseUDP(pkt)
|
||||||
|
if !ok {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return c.commitParsed(pkt, info)
|
||||||
|
}
|
||||||
|
|
||||||
|
// commitParsed is the post-parse half of Commit. The caller must have
|
||||||
|
// already verified parseUDP succeeded. Used by MultiCoalescer.Commit to
|
||||||
|
// avoid re-walking the IP/UDP header.
|
||||||
|
func (c *UDPCoalescer) commitParsed(pkt []byte, info parsedUDP) error {
|
||||||
|
if c.gsoW == nil {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
// A zero-length UDP datagram (UDP `length` == 8) is legal and must still
|
||||||
|
// reach the TUN, but it can't be coalesced: a GSO slot would store an
|
||||||
|
// empty payload iovec and the kernel has nothing to segment. Seal any
|
||||||
|
// open chain for this flow (so a later, non-empty datagram seeds fresh
|
||||||
|
// *after* this one and per-flow arrival order is preserved) and deliver
|
||||||
|
// it as a plain single datagram.
|
||||||
|
if info.payLen == 0 {
|
||||||
|
delete(c.openSlots, info.fk)
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if open := c.openSlots[info.fk]; open != nil {
|
||||||
|
if c.canAppend(open, pkt, info) {
|
||||||
|
c.appendPayload(open, pkt, info)
|
||||||
|
if open.sealed {
|
||||||
|
delete(c.openSlots, info.fk)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
// Can't extend — seal it and fall through to seed a fresh slot.
|
||||||
|
delete(c.openSlots, info.fk)
|
||||||
|
}
|
||||||
|
c.seed(pkt, info)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Flush drains every queued slot and calls the configured Resetter.
|
||||||
|
func (c *UDPCoalescer) Flush() error {
|
||||||
|
first := c.drain()
|
||||||
|
if c.resetter != nil {
|
||||||
|
c.resetter()
|
||||||
|
}
|
||||||
|
return first
|
||||||
|
}
|
||||||
|
|
||||||
|
// drain emits every queued slot in arrival order and clears the slot state.
|
||||||
|
// It does NOT reset the arena: borrowed payload slices stay valid until the
|
||||||
|
// arena's owner recycles it.
|
||||||
|
func (c *UDPCoalescer) drain() error {
|
||||||
|
var first error
|
||||||
|
for _, s := range c.slots {
|
||||||
|
var err error
|
||||||
|
if s.passthrough {
|
||||||
|
_, err = c.plainW.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)
|
||||||
|
return first
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCoalescer) addPassthrough(pkt []byte) {
|
||||||
|
s := c.take()
|
||||||
|
s.passthrough = true
|
||||||
|
s.rawPkt = pkt
|
||||||
|
c.slots = append(c.slots, s)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCoalescer) seed(pkt []byte, info parsedUDP) {
|
||||||
|
if info.hdrLen > udpCoalesceHdrCap || info.hdrLen+info.payLen > udpCoalesceBufSize {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s := c.take()
|
||||||
|
s.passthrough = false
|
||||||
|
s.rawPkt = nil
|
||||||
|
copy(s.hdrBuf[:], pkt[:info.hdrLen])
|
||||||
|
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.sealed = false
|
||||||
|
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||||
|
c.slots = append(c.slots, s)
|
||||||
|
c.openSlots[info.fk] = s
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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 s.sealed {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
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
|
||||||
|
}
|
||||||
|
if !udpHeadersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCoalescer) appendPayload(s *udpSlot, pkt []byte, info parsedUDP) {
|
||||||
|
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||||
|
s.numSeg++
|
||||||
|
s.totalPay += info.payLen
|
||||||
|
if info.payLen < s.gsoSize {
|
||||||
|
// Last-segment-can-be-shorter: this seals the chain.
|
||||||
|
s.sealed = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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) {
|
||||||
|
s.passthrough = false
|
||||||
|
s.rawPkt = nil
|
||||||
|
clear(s.payIovs)
|
||||||
|
s.payIovs = s.payIovs[:0]
|
||||||
|
s.numSeg = 0
|
||||||
|
s.totalPay = 0
|
||||||
|
s.sealed = false
|
||||||
|
c.pool = append(c.pool, s)
|
||||||
|
}
|
||||||
|
|
||||||
|
// flushSlot patches the IP header total length / IPv6 payload length and
|
||||||
|
// the UDP length to the *total* across all coalesced segments, then seeds
|
||||||
|
// the UDP checksum field with the pseudo-header partial (single-fold, not
|
||||||
|
// inverted) per virtio NEEDS_CSUM. The kernel's ip_rcv_core (v4) and
|
||||||
|
// ip6_rcv_core (v6) trim the skb to those length fields, so per-segment
|
||||||
|
// values would silently drop everything but the first segment. The kernel
|
||||||
|
// then walks each segment in __udp_gso_segment, recomputing per-segment
|
||||||
|
// uh->len / iph->tot_len / IPv6 plen and adjusting the checksum via
|
||||||
|
// `check = csum16_add(csum16_sub(uh->check, uh->len), newlen)` — meaning
|
||||||
|
// our seed's uh->check must be consistent with the seed's uh->len, which
|
||||||
|
// is what passing the total to both pseudoSum and the UDP length field
|
||||||
|
// guarantees.
|
||||||
|
func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
|
||||||
|
hdr := s.hdrBuf[: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.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoUDP)
|
||||||
|
}
|
||||||
|
|
||||||
|
// udpHeadersMatch compares two IP+UDP header prefixes for byte-equality on
|
||||||
|
// every field that must be identical across coalesced segments. Length
|
||||||
|
// fields are masked out (flushSlot rewrites them), but the IP-level ECN
|
||||||
|
// codepoint is compared (via ipHeadersMatch) so segments with differing ECN
|
||||||
|
// don't coalesce, matching kernel GRO.
|
||||||
|
func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
||||||
|
if len(a) != len(b) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
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
|
||||||
|
if a[udp] != b[udp] || a[udp+1] != b[udp+1] || a[udp+2] != b[udp+2] || a[udp+3] != b[udp+3] {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
@@ -0,0 +1,472 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"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
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUDPCoalescerPassthroughWhenGSOUnavailable(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: false}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
pkt := buildUDPv4(1000, 53, make([]byte, 100))
|
||||||
|
if err := c.Commit(pkt); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(w.writes) != 0 || len(w.gsoWrites) != 0 {
|
||||||
|
t.Fatalf("no Add-time writes: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
if err := c.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||||
|
t.Fatalf("want single plain write, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
// 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}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
pkt := buildUDPv4(1000, 53, make([]byte, 800))
|
||||||
|
if err := c.Commit(pkt); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := c.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
// Single-segment flush goes through WriteGSO; the writer infers GSO_NONE
|
||||||
|
// from len(pays)==1 and the kernel fills in the UDP csum (NEEDS_CSUM).
|
||||||
|
if len(w.gsoWrites) != 1 || len(w.writes) != 0 {
|
||||||
|
t.Fatalf("single-seg flush: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
pay := make([]byte, 1200)
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||||
|
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}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
full := make([]byte, 1200)
|
||||||
|
tail := make([]byte, 600)
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites) != 2 {
|
||||||
|
t.Fatalf("want 2 gso writes (sealed + new seed), got %d", len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites[0].pays) != 3 {
|
||||||
|
t.Errorf("first super: want 3 pays, got %d", len(w.gsoWrites[0].pays))
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites[1].pays) != 1 {
|
||||||
|
t.Errorf("second super: want 1 pay (re-seed), got %d", len(w.gsoWrites[1].pays))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A larger-than-gsoSize packet cannot extend the slot — it reseeds.
|
||||||
|
func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 1200))); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := c.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites) != 2 {
|
||||||
|
t.Fatalf("want 2 separate seeds, got %d", len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Different 5-tuples must not coalesce.
|
||||||
|
func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
pay := make([]byte, 800)
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
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}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
pay := make([]byte, 100)
|
||||||
|
for i := 0; i < udpCoalesceMaxSegs+5; i++ {
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||||
|
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 seeds a fresh superpacket that keeps CE; the
|
||||||
|
// trailing Not-ECT datagram seeds another.
|
||||||
|
func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
pay := make([]byte, 800)
|
||||||
|
pkt0 := buildUDPv4(1000, 53, pay) // ECN=00 (Not-ECT)
|
||||||
|
pkt1 := buildUDPv4(1000, 53, pay)
|
||||||
|
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.gsoWrites) != 3 {
|
||||||
|
t.Fatalf("want 3 separate seeds (differing ECN), got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
||||||
|
}
|
||||||
|
wantECN := []byte{0x00, 0x03, 0x00}
|
||||||
|
for i, g := range w.gsoWrites {
|
||||||
|
if len(g.pays) != 1 {
|
||||||
|
t.Errorf("gso %d pay count=%d want 1", i, len(g.pays))
|
||||||
|
}
|
||||||
|
if got := g.hdr[1] & 0x03; got != wantECN[i] {
|
||||||
|
t.Errorf("gso %d ECN=%#x want %#x", i, got, wantECN[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// IPv6 path: same flow, equal-sized → coalesced.
|
||||||
|
func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
pay := make([]byte, 1200)
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
|
||||||
|
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}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites) != 2 {
|
||||||
|
t.Fatalf("want 2 separate seeds (different DSCP), got %d", len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fragmented IPv4 must not be coalesced.
|
||||||
|
func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
||||||
|
binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1
|
||||||
|
if err := c.Commit(pkt); err != nil {
|
||||||
|
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}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
pkt := buildUDPv4(1000, 53, nil) // UDP length 8, zero payload
|
||||||
|
if err := c.Commit(pkt); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
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 passthrough contract as v4.
|
||||||
|
func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
pkt := buildUDPv6(1000, 53, nil) // UDP length 8, zero payload
|
||||||
|
if err := c.Commit(pkt); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
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}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
full := make([]byte, 800)
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
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: two single-segment superpackets bracket one plain write.
|
||||||
|
if len(w.gsoWrites) != 2 || len(w.writes) != 1 {
|
||||||
|
t.Fatalf("want 2 gso writes + 1 plain, got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// IPv4 with options is not admissible (we require IHL=5).
|
||||||
|
func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
arena := NewArena(0)
|
||||||
|
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
|
||||||
|
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
||||||
|
pkt[0] = 0x46 // IHL = 6 (24-byte IPv4 header — has options)
|
||||||
|
if err := c.Commit(pkt); err != nil {
|
||||||
|
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))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,23 @@
|
|||||||
|
package checksum
|
||||||
|
|
||||||
|
import (
|
||||||
|
"golang.org/x/sys/cpu"
|
||||||
|
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||||
|
)
|
||||||
|
|
||||||
|
//go:noescape
|
||||||
|
func checksumAVX2(buf []byte, initial uint16) uint16
|
||||||
|
|
||||||
|
var hasAVX2 = cpu.X86.HasAVX2
|
||||||
|
|
||||||
|
// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with
|
||||||
|
// initial. It is a drop-in replacement for gvisor's checksum.Checksum that
|
||||||
|
// dispatches to a hand-written AVX2 routine on amd64 CPUs that support it,
|
||||||
|
// falling back to gvisor's pure-Go implementation otherwise. The result
|
||||||
|
// matches gvisor's bit-for-bit for any buffer length and initial seed.
|
||||||
|
func Checksum(buf []byte, initial uint16) uint16 {
|
||||||
|
if hasAVX2 {
|
||||||
|
return checksumAVX2(buf, initial)
|
||||||
|
}
|
||||||
|
return gvisorchecksum.Checksum(buf, initial)
|
||||||
|
}
|
||||||
@@ -0,0 +1,157 @@
|
|||||||
|
#include "textflag.h"
|
||||||
|
|
||||||
|
// func checksumAVX2(buf []byte, initial uint16) uint16
|
||||||
|
//
|
||||||
|
// Computes the RFC 1071 ones-complement sum of buf, seeded with initial.
|
||||||
|
//
|
||||||
|
// Algorithm: sum the buffer treating it as a stream of uint32s in machine
|
||||||
|
// (little-endian) byte order, accumulating into 64-bit lanes (top 32 bits
|
||||||
|
// hold cross-add carries — at 1 byte / lane / iter we have 32 bits of
|
||||||
|
// headroom which is far more than the 16 KB/64 KB max practical inputs).
|
||||||
|
// At the end we fold to 16 bits and byte-swap once to recover the on-wire
|
||||||
|
// (big-endian) result. RFC 1071 §1.2.B byte-order independence makes this
|
||||||
|
// equivalent to summing as 16-bit big-endian words.
|
||||||
|
//
|
||||||
|
// The ymm accumulators (Y4..Y7) hold 4 uint64 lanes each = 16 parallel
|
||||||
|
// partial sums. The main loop loads 64 bytes per iter as four 16-byte
|
||||||
|
// chunks, zero-extending each chunk's four uint32s into a ymm via
|
||||||
|
// VPMOVZXDQ-from-memory, then VPADDQ into a separate accumulator per
|
||||||
|
// chunk to break the dep chain. After the vector loop the lane sums are
|
||||||
|
// horizontally reduced and merged with a scalar accumulator that handles
|
||||||
|
// the trailing 0..63 bytes plus the (byte-swapped) initial seed.
|
||||||
|
TEXT ·checksumAVX2(SB), NOSPLIT, $0-34
|
||||||
|
MOVQ buf_base+0(FP), SI
|
||||||
|
MOVQ buf_len+8(FP), CX
|
||||||
|
MOVWQZX initial+24(FP), AX
|
||||||
|
|
||||||
|
// Pre-byteswap initial into the LE-summing space so it merges directly
|
||||||
|
// with the rest of the accumulator. The final fold's bswap16 will undo
|
||||||
|
// this and convert the whole result back to BE.
|
||||||
|
XCHGB AH, AL
|
||||||
|
|
||||||
|
CMPQ CX, $32
|
||||||
|
JLT scalar_tail
|
||||||
|
|
||||||
|
VPXOR Y4, Y4, Y4
|
||||||
|
VPXOR Y5, Y5, Y5
|
||||||
|
VPXOR Y6, Y6, Y6
|
||||||
|
VPXOR Y7, Y7, Y7
|
||||||
|
|
||||||
|
CMPQ CX, $64
|
||||||
|
JLT loop32
|
||||||
|
|
||||||
|
loop64:
|
||||||
|
VPMOVZXDQ (SI), Y0
|
||||||
|
VPMOVZXDQ 16(SI), Y1
|
||||||
|
VPMOVZXDQ 32(SI), Y2
|
||||||
|
VPMOVZXDQ 48(SI), Y3
|
||||||
|
VPADDQ Y0, Y4, Y4
|
||||||
|
VPADDQ Y1, Y5, Y5
|
||||||
|
VPADDQ Y2, Y6, Y6
|
||||||
|
VPADDQ Y3, Y7, Y7
|
||||||
|
ADDQ $64, SI
|
||||||
|
SUBQ $64, CX
|
||||||
|
CMPQ CX, $64
|
||||||
|
JGE loop64
|
||||||
|
|
||||||
|
loop32:
|
||||||
|
CMPQ CX, $32
|
||||||
|
JLT reduce_vec
|
||||||
|
VPMOVZXDQ (SI), Y0
|
||||||
|
VPMOVZXDQ 16(SI), Y1
|
||||||
|
VPADDQ Y0, Y4, Y4
|
||||||
|
VPADDQ Y1, Y5, Y5
|
||||||
|
ADDQ $32, SI
|
||||||
|
SUBQ $32, CX
|
||||||
|
JMP loop32
|
||||||
|
|
||||||
|
reduce_vec:
|
||||||
|
// Combine the four ymm accumulators into Y4.
|
||||||
|
VPADDQ Y5, Y4, Y4
|
||||||
|
VPADDQ Y7, Y6, Y6
|
||||||
|
VPADDQ Y6, Y4, Y4
|
||||||
|
|
||||||
|
// Horizontally reduce Y4's four uint64 lanes to a single scalar.
|
||||||
|
VEXTRACTI128 $1, Y4, X5
|
||||||
|
VPADDQ X5, X4, X4
|
||||||
|
VPSHUFD $0x4e, X4, X5
|
||||||
|
VPADDQ X5, X4, X4
|
||||||
|
VMOVQ X4, R8
|
||||||
|
VZEROUPPER
|
||||||
|
|
||||||
|
ADDQ R8, AX
|
||||||
|
ADCQ $0, AX
|
||||||
|
|
||||||
|
scalar_tail:
|
||||||
|
// Handle remaining 0..63 bytes (or the entire buffer if it was < 32).
|
||||||
|
CMPQ CX, $8
|
||||||
|
JLT tail4
|
||||||
|
|
||||||
|
loop8:
|
||||||
|
ADDQ (SI), AX
|
||||||
|
ADCQ $0, AX
|
||||||
|
ADDQ $8, SI
|
||||||
|
SUBQ $8, CX
|
||||||
|
CMPQ CX, $8
|
||||||
|
JGE loop8
|
||||||
|
|
||||||
|
tail4:
|
||||||
|
CMPQ CX, $4
|
||||||
|
JLT tail2
|
||||||
|
MOVL (SI), R8
|
||||||
|
ADDQ R8, AX
|
||||||
|
ADCQ $0, AX
|
||||||
|
ADDQ $4, SI
|
||||||
|
SUBQ $4, CX
|
||||||
|
|
||||||
|
tail2:
|
||||||
|
CMPQ CX, $2
|
||||||
|
JLT tail1
|
||||||
|
MOVWQZX (SI), R8
|
||||||
|
ADDQ R8, AX
|
||||||
|
ADCQ $0, AX
|
||||||
|
ADDQ $2, SI
|
||||||
|
SUBQ $2, CX
|
||||||
|
|
||||||
|
tail1:
|
||||||
|
TESTQ CX, CX
|
||||||
|
JZ fold
|
||||||
|
MOVBQZX (SI), R8
|
||||||
|
ADDQ R8, AX
|
||||||
|
ADCQ $0, AX
|
||||||
|
|
||||||
|
fold:
|
||||||
|
// Fold the 64-bit accumulator to 16 bits via four rounds, mirroring
|
||||||
|
// gvisor's reduce(). Each pair (split, add) halves the live width;
|
||||||
|
// the truncation steps absorb the single bit that may be left over
|
||||||
|
// after each add so the next round's bound holds.
|
||||||
|
|
||||||
|
// 64 → 33 bits.
|
||||||
|
MOVQ AX, R8
|
||||||
|
SHRQ $32, R8
|
||||||
|
MOVL AX, AX
|
||||||
|
ADDQ R8, AX
|
||||||
|
|
||||||
|
// 33 → 32 bits. AX += (AX>>32); truncate to 32. AX is now ≤ 0xFFFF_FFFF.
|
||||||
|
MOVQ AX, R8
|
||||||
|
SHRQ $32, R8
|
||||||
|
ADDQ R8, AX
|
||||||
|
MOVL AX, AX
|
||||||
|
|
||||||
|
// 32 → 17 bits.
|
||||||
|
MOVQ AX, R8
|
||||||
|
SHRQ $16, R8
|
||||||
|
MOVWQZX AX, AX
|
||||||
|
ADDQ R8, AX
|
||||||
|
|
||||||
|
// 17 → 16 bits. AX += (AX>>16); the trailing MOVW truncates bit 16.
|
||||||
|
MOVQ AX, R8
|
||||||
|
SHRQ $16, R8
|
||||||
|
ADDQ R8, AX
|
||||||
|
|
||||||
|
// AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip
|
||||||
|
// to big-endian to match the gvisor API contract.
|
||||||
|
XCHGB AH, AL
|
||||||
|
|
||||||
|
MOVW AX, ret+32(FP)
|
||||||
|
RET
|
||||||
@@ -0,0 +1,12 @@
|
|||||||
|
package checksum
|
||||||
|
|
||||||
|
//go:noescape
|
||||||
|
func checksumNEON(buf []byte, initial uint16) uint16
|
||||||
|
|
||||||
|
// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with
|
||||||
|
// initial. It is a drop-in replacement for gvisor's checksum.Checksum
|
||||||
|
// that dispatches to a hand-written NEON routine. NEON is mandatory in
|
||||||
|
// armv8 so no feature check is needed.
|
||||||
|
func Checksum(buf []byte, initial uint16) uint16 {
|
||||||
|
return checksumNEON(buf, initial)
|
||||||
|
}
|
||||||
@@ -0,0 +1,143 @@
|
|||||||
|
#include "textflag.h"
|
||||||
|
|
||||||
|
// func checksumNEON(buf []byte, initial uint16) uint16
|
||||||
|
//
|
||||||
|
// Mirrors the algorithm in checksum_amd64.s: sum the buffer treating it as
|
||||||
|
// a stream of uint32s in machine (little-endian) byte order, accumulating
|
||||||
|
// into 64-bit lanes that have ample carry headroom; fold and byte-swap once
|
||||||
|
// at the very end to recover the on-wire (big-endian) result.
|
||||||
|
//
|
||||||
|
// Each loop iteration loads 64 bytes via VLD1.P into V0..V3 (4 Q regs).
|
||||||
|
// VUADDW takes the low two uint32 lanes of a Q reg, zero-extends them to
|
||||||
|
// uint64, and adds them into a 2×uint64 accumulator; VUADDW2 does the same
|
||||||
|
// for the high two lanes. Four ymm-equivalent accumulators (V8..V11) get
|
||||||
|
// updated twice per iter to break the dep chain. Tail bytes go through a
|
||||||
|
// scalar ADCS chain seeded with the byte-swapped initial.
|
||||||
|
TEXT ·checksumNEON(SB), NOSPLIT, $0-34
|
||||||
|
MOVD buf_base+0(FP), R0
|
||||||
|
MOVD buf_len+8(FP), R1
|
||||||
|
MOVHU initial+24(FP), R2
|
||||||
|
|
||||||
|
// Pre-byteswap initial into the LE-summing space so it merges directly
|
||||||
|
// with the rest of the accumulator.
|
||||||
|
REV16W R2, R2
|
||||||
|
|
||||||
|
MOVD ZR, R3 // scalar accumulator
|
||||||
|
|
||||||
|
CMP $32, R1
|
||||||
|
BLT scalar_tail
|
||||||
|
|
||||||
|
VEOR V8.B16, V8.B16, V8.B16
|
||||||
|
VEOR V9.B16, V9.B16, V9.B16
|
||||||
|
VEOR V10.B16, V10.B16, V10.B16
|
||||||
|
VEOR V11.B16, V11.B16, V11.B16
|
||||||
|
|
||||||
|
CMP $64, R1
|
||||||
|
BLT loop16_init
|
||||||
|
|
||||||
|
loop64:
|
||||||
|
VLD1.P 64(R0), [V0.B16, V1.B16, V2.B16, V3.B16]
|
||||||
|
VUADDW V0.S2, V8.D2, V8.D2
|
||||||
|
VUADDW2 V0.S4, V9.D2, V9.D2
|
||||||
|
VUADDW V1.S2, V10.D2, V10.D2
|
||||||
|
VUADDW2 V1.S4, V11.D2, V11.D2
|
||||||
|
VUADDW V2.S2, V8.D2, V8.D2
|
||||||
|
VUADDW2 V2.S4, V9.D2, V9.D2
|
||||||
|
VUADDW V3.S2, V10.D2, V10.D2
|
||||||
|
VUADDW2 V3.S4, V11.D2, V11.D2
|
||||||
|
SUB $64, R1, R1
|
||||||
|
CMP $64, R1
|
||||||
|
BGE loop64
|
||||||
|
|
||||||
|
loop16_init:
|
||||||
|
CMP $16, R1
|
||||||
|
BLT reduce_vec
|
||||||
|
|
||||||
|
loop16:
|
||||||
|
VLD1.P 16(R0), [V0.B16]
|
||||||
|
VUADDW V0.S2, V8.D2, V8.D2
|
||||||
|
VUADDW2 V0.S4, V9.D2, V9.D2
|
||||||
|
SUB $16, R1, R1
|
||||||
|
CMP $16, R1
|
||||||
|
BGE loop16
|
||||||
|
|
||||||
|
reduce_vec:
|
||||||
|
// Combine the four accumulators into V8.
|
||||||
|
VADD V9.D2, V8.D2, V8.D2
|
||||||
|
VADD V11.D2, V10.D2, V10.D2
|
||||||
|
VADD V10.D2, V8.D2, V8.D2
|
||||||
|
|
||||||
|
// Horizontal-add the two lanes of V8.D2 into a single uint64.
|
||||||
|
VADDP V8.D2, V8.D2, V8.D2
|
||||||
|
VMOV V8.D[0], R8
|
||||||
|
|
||||||
|
ADDS R8, R3, R3
|
||||||
|
ADC ZR, R3, R3
|
||||||
|
|
||||||
|
scalar_tail:
|
||||||
|
CMP $8, R1
|
||||||
|
BLT tail4
|
||||||
|
|
||||||
|
loop8:
|
||||||
|
MOVD.P 8(R0), R8
|
||||||
|
ADDS R8, R3, R3
|
||||||
|
ADC ZR, R3, R3
|
||||||
|
SUB $8, R1, R1
|
||||||
|
CMP $8, R1
|
||||||
|
BGE loop8
|
||||||
|
|
||||||
|
tail4:
|
||||||
|
CMP $4, R1
|
||||||
|
BLT tail2
|
||||||
|
MOVWU.P 4(R0), R8
|
||||||
|
ADDS R8, R3, R3
|
||||||
|
ADC ZR, R3, R3
|
||||||
|
SUB $4, R1, R1
|
||||||
|
|
||||||
|
tail2:
|
||||||
|
CMP $2, R1
|
||||||
|
BLT tail1
|
||||||
|
MOVHU.P 2(R0), R8
|
||||||
|
ADDS R8, R3, R3
|
||||||
|
ADC ZR, R3, R3
|
||||||
|
SUB $2, R1, R1
|
||||||
|
|
||||||
|
tail1:
|
||||||
|
CBZ R1, fold
|
||||||
|
MOVBU (R0), R8
|
||||||
|
ADDS R8, R3, R3
|
||||||
|
ADC ZR, R3, R3
|
||||||
|
|
||||||
|
fold:
|
||||||
|
// Merge the byte-swapped initial into our LE-form accumulator.
|
||||||
|
ADDS R2, R3, R3
|
||||||
|
ADC ZR, R3, R3
|
||||||
|
|
||||||
|
// 64 → 33 bits.
|
||||||
|
LSR $32, R3, R8
|
||||||
|
AND $0xffffffff, R3, R3
|
||||||
|
ADD R8, R3, R3
|
||||||
|
|
||||||
|
// 33 → 32 (truncate after adding bit 32 back).
|
||||||
|
LSR $32, R3, R8
|
||||||
|
ADD R8, R3, R3
|
||||||
|
AND $0xffffffff, R3, R3
|
||||||
|
|
||||||
|
// 32 → 17.
|
||||||
|
LSR $16, R3, R8
|
||||||
|
AND $0xffff, R3, R3
|
||||||
|
ADD R8, R3, R3
|
||||||
|
|
||||||
|
// 17 → 16 (truncation absorbs bit 16 below).
|
||||||
|
LSR $16, R3, R8
|
||||||
|
ADD R8, R3, R3
|
||||||
|
|
||||||
|
// AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip
|
||||||
|
// to big-endian to match the gvisor API contract. REV16W swaps bytes
|
||||||
|
// within each 16-bit halfword of the low 32 bits, so it acts as a
|
||||||
|
// 16-bit byte-swap on the live low 16.
|
||||||
|
REV16W R3, R3
|
||||||
|
AND $0xffff, R3, R3
|
||||||
|
|
||||||
|
MOVH R3, ret+32(FP)
|
||||||
|
RET
|
||||||
@@ -0,0 +1,10 @@
|
|||||||
|
//go:build !amd64 && !arm64
|
||||||
|
|
||||||
|
package checksum
|
||||||
|
|
||||||
|
import gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||||
|
|
||||||
|
// Checksum delegates to gvisor on architectures without a hand-written body.
|
||||||
|
func Checksum(buf []byte, initial uint16) uint16 {
|
||||||
|
return gvisorchecksum.Checksum(buf, initial)
|
||||||
|
}
|
||||||
@@ -0,0 +1,190 @@
|
|||||||
|
package checksum
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"math/rand/v2"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestChecksumMatchesGvisor walks lengths from 0 to 4096, with several initial
|
||||||
|
// seeds and a handful of starting alignments, asserting that our local
|
||||||
|
// Checksum matches gvisor's reference bit-for-bit.
|
||||||
|
func TestChecksumMatchesGvisor(t *testing.T) {
|
||||||
|
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 := Checksum(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 length := 0; length <= 256; length++ {
|
||||||
|
patterns := map[string][]byte{
|
||||||
|
"zeros": make([]byte, length),
|
||||||
|
"ones": bytes(length, 0xff),
|
||||||
|
"alternating": pattern(length, []byte{0xa5, 0x5a}),
|
||||||
|
"ascending": ascending(length),
|
||||||
|
}
|
||||||
|
for name, buf := range patterns {
|
||||||
|
for _, seed := range []uint16{0, 0xffff, 0x8000} {
|
||||||
|
want := gvisorchecksum.Checksum(buf, seed)
|
||||||
|
got := Checksum(buf, seed)
|
||||||
|
if got != want {
|
||||||
|
t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x",
|
||||||
|
name, length, seed, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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) {
|
||||||
|
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 := Checksum(buf, seed)
|
||||||
|
if got != want {
|
||||||
|
t.Fatalf("k=%d tail=%d (len=%d) off=%d seed=%#x: got %#04x want %#04x",
|
||||||
|
k, tail, length, off, seed, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
+2
-2
@@ -8,8 +8,8 @@ import (
|
|||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
// defaultBatchBufSize is the per-Queue scratch size for Read. 65535 covers
|
// defaultBatchBufSize is the per-Queue scratch size for Read on backends
|
||||||
// any single IP packet.
|
// that don't do TSO segmentation. 65535 covers any single IP packet.
|
||||||
const defaultBatchBufSize = 65535
|
const defaultBatchBufSize = 65535
|
||||||
|
|
||||||
type Device interface {
|
type Device interface {
|
||||||
|
|||||||
@@ -0,0 +1,99 @@
|
|||||||
|
//go:build linux && !android
|
||||||
|
// +build linux,!android
|
||||||
|
|
||||||
|
package tio
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"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
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do
|
||||||
|
// TSO segmentation in userspace. usoEnabled tells downstream queues whether
|
||||||
|
// the kernel agreed to deliver/accept GSO_UDP_L4 superpackets — coalescers
|
||||||
|
// should fall back to per-packet writes when this is false.
|
||||||
|
func NewOffloadQueueSet(usoEnabled bool) (QueueSet, error) {
|
||||||
|
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
out := &offloadQueueSet{
|
||||||
|
pq: []*Offload{},
|
||||||
|
pqi: []Queue{},
|
||||||
|
shutdownFd: shutdownFd,
|
||||||
|
usoEnabled: usoEnabled,
|
||||||
|
}
|
||||||
|
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *offloadQueueSet) Queues() []Queue {
|
||||||
|
return c.pqi
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *offloadQueueSet) Add(fd int) error {
|
||||||
|
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled)
|
||||||
|
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.
|
||||||
|
// The per-queue Close deliberately leaves shutdownFd alone - it belongs
|
||||||
|
// to this container.
|
||||||
|
for _, x := range c.pq {
|
||||||
|
if err := x.Close(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close the shutdown eventfd last: every reader's pollfd set references
|
||||||
|
// it, so it must outlive the wake + per-queue teardown above.
|
||||||
|
if err := unix.Close(c.shutdownFd); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
c.shutdownFd = -1
|
||||||
|
|
||||||
|
return errors.Join(errs...)
|
||||||
|
}
|
||||||
@@ -0,0 +1,65 @@
|
|||||||
|
//go:build linux && !android && !e2e_testing
|
||||||
|
|
||||||
|
package tio
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
// fakeBatch stands in for batch.TxBatcher inside the bench — same shape
|
||||||
|
// of pointer-capturing closure that sendInsideMessage builds.
|
||||||
|
type fakeBatch struct{ buf [65536]byte }
|
||||||
|
|
||||||
|
func (b *fakeBatch) Reserve(sz int) []byte { return b.buf[:sz] }
|
||||||
|
func (b *fakeBatch) Commit([]byte) {}
|
||||||
|
|
||||||
|
type fakeHostInfo struct {
|
||||||
|
remoteIndexId uint32
|
||||||
|
counter uint64
|
||||||
|
}
|
||||||
|
type fakeIface struct {
|
||||||
|
rebindCount uint8
|
||||||
|
hi *fakeHostInfo
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkSegmentSuperpacketAllocsTSO measures allocation per
|
||||||
|
// SegmentSuperpacket call when a closure captures pointer-bearing
|
||||||
|
// receivers — the realistic shape of sendInsideMessage's closure.
|
||||||
|
func BenchmarkSegmentSuperpacketAllocsTSO(b *testing.B) {
|
||||||
|
const mss = 1400
|
||||||
|
const numSeg = 32
|
||||||
|
pkt := buildTSOv6(mss*numSeg, mss)
|
||||||
|
gso := GSOInfo{
|
||||||
|
Size: mss,
|
||||||
|
HdrLen: 60, // 40 (IPv6) + 20 (TCP)
|
||||||
|
CsumStart: 40,
|
||||||
|
Proto: GSOProtoTCP,
|
||||||
|
}
|
||||||
|
p := Packet{Bytes: pkt, GSO: gso}
|
||||||
|
|
||||||
|
hi := &fakeHostInfo{remoteIndexId: 0xdeadbeef}
|
||||||
|
f := &fakeIface{rebindCount: 7, hi: hi}
|
||||||
|
fb := &fakeBatch{}
|
||||||
|
|
||||||
|
// SegmentSuperpacket consumes pkt destructively; refresh from a master
|
||||||
|
// copy each iter (matches the production pattern where every TUN read
|
||||||
|
// hands the segmenter a fresh kernel-supplied buffer).
|
||||||
|
master := append([]byte(nil), pkt...)
|
||||||
|
work := make([]byte, len(pkt))
|
||||||
|
p.Bytes = work
|
||||||
|
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
copy(work, master)
|
||||||
|
err := SegmentSuperpacket(p, func(seg []byte) error {
|
||||||
|
out := fb.Reserve(16 + len(seg) + 16)
|
||||||
|
out[0] = byte(f.rebindCount)
|
||||||
|
out[1] = byte(hi.counter)
|
||||||
|
hi.counter++
|
||||||
|
fb.Commit(out)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
b.Fatalf("SegmentSuperpacket: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
//go:build !linux || android || e2e_testing
|
||||||
|
|
||||||
|
package tio
|
||||||
|
|
||||||
|
import "fmt"
|
||||||
|
|
||||||
|
func protoFromGSOType(_ uint8) (GSOProto, error) {
|
||||||
|
return 0, fmt.Errorf("GSO unsupported")
|
||||||
|
}
|
||||||
|
|
||||||
|
// SegmentSuperpacket invokes fn once per segment of pkt. On non-Linux
|
||||||
|
// builds (and Android/e2e_testing) this package does not provide a Queue
|
||||||
|
// implementation, so any caller that does construct a Packet here can only
|
||||||
|
// be operating on non-superpacket bytes and the stub forwards them
|
||||||
|
// directly. A non-zero GSO field is a programming error from the caller
|
||||||
|
// and returns an explicit error rather than silently misbehaving.
|
||||||
|
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
||||||
|
if pkt.GSO.IsSuperpacket() {
|
||||||
|
return fmt.Errorf("tio: GSO superpacket on platform without segmentation support")
|
||||||
|
}
|
||||||
|
return fn(pkt.Bytes)
|
||||||
|
}
|
||||||
+129
-9
@@ -13,16 +13,33 @@ type QueueSet interface {
|
|||||||
Add(fd int) error
|
Add(fd int) error
|
||||||
}
|
}
|
||||||
|
|
||||||
// Queue is a readable/writable packet queue. Concurrency contract: a single
|
// Capabilities advertises which kernel offload features a Queue
|
||||||
// read goroutine drives Read; plain Write is safe for concurrent callers.
|
// successfully negotiated. Callers consult this to decide which coalescers
|
||||||
|
// to wire onto the write path — a Queue without TSO can't usefully accept a
|
||||||
|
// TCPCoalescer, and a Queue without USO can't accept a UDPCoalescer.
|
||||||
|
type Capabilities struct {
|
||||||
|
// TSO means the FD was opened with IFF_VNET_HDR and the kernel agreed
|
||||||
|
// to TUN_F_TSO4|TSO6 — i.e. WriteGSO with GSOProtoTCP is safe.
|
||||||
|
TSO bool
|
||||||
|
// USO means the kernel additionally agreed to TUN_F_USO4|USO6, so
|
||||||
|
// WriteGSO with GSOProtoUDP is safe. Linux ≥ 6.2.
|
||||||
|
USO 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.
|
||||||
type Queue interface {
|
type Queue interface {
|
||||||
io.Closer
|
io.Closer
|
||||||
|
|
||||||
// Read returns one or more packets. The returned Packet.Bytes slices
|
// Read returns one or more packets. The returned Packet.Bytes slices
|
||||||
// are borrowed from the Queue's internal buffer and are only valid
|
// are borrowed from the Queue's internal buffer and are only valid
|
||||||
// until the next Read or Close on this Queue - callers must encrypt
|
// until the next Read or Close on this Queue - callers must encrypt
|
||||||
// or copy each slice before the next call. Single-reader only: not
|
// or copy each slice before the next call. A Packet may carry a
|
||||||
// safe for concurrent Reads (it reuses per-queue rx scratch each call).
|
// GSO/USO superpacket (see GSOInfo); when GSO.IsSuperpacket() is
|
||||||
|
// true the caller must segment Bytes before treating it as a single
|
||||||
|
// IP datagram. Single-reader only: not safe for concurrent Reads (it
|
||||||
|
// reuses per-queue rx scratch each call).
|
||||||
Read() ([]Packet, error)
|
Read() ([]Packet, error)
|
||||||
|
|
||||||
// Write emits a single packet on the plaintext (outside→inside)
|
// Write emits a single packet on the plaintext (outside→inside)
|
||||||
@@ -32,21 +49,124 @@ type Queue interface {
|
|||||||
|
|
||||||
// Packet is the unit Queue.Read returns. Bytes points into the queue's
|
// 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
|
// internal buffer and is only valid until the next Read or Close on the
|
||||||
// queue that produced it.
|
// queue that produced it. GSO is the zero value for an already-segmented
|
||||||
|
// IP datagram; when non-zero it describes a kernel-supplied TSO/USO
|
||||||
|
// superpacket the caller must segment before consuming.
|
||||||
type Packet struct {
|
type Packet struct {
|
||||||
Bytes []byte
|
Bytes []byte
|
||||||
|
GSO GSOInfo
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GSOInfo describes a kernel-supplied superpacket sitting in Packet.Bytes.
|
||||||
|
// The zero value means "not a superpacket" — Bytes is one regular IP
|
||||||
|
// datagram and no segmentation is required.
|
||||||
|
type GSOInfo struct {
|
||||||
|
// Size is the GSO segment size: max payload bytes per segment
|
||||||
|
// (== TCP MSS for TSO, == UDP payload chunk for USO). Zero means
|
||||||
|
// not a superpacket.
|
||||||
|
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,
|
// Clone returns a Packet whose Bytes is a freshly allocated copy of p.Bytes,
|
||||||
// safe to retain past the next Read or Close on the originating Queue.
|
// safe to retain past the next Read or Close on the originating Queue.
|
||||||
// Use this only when a caller genuinely needs to outlive the borrowed-slice
|
// GSO metadata is copied verbatim. Use this only when a caller genuinely
|
||||||
// contract — the hot path reads should continue to consume the borrow
|
// needs to outlive the borrowed-slice contract — the hot path reads should
|
||||||
// synchronously to avoid the allocation.
|
// continue to consume the borrow synchronously to avoid the allocation.
|
||||||
func (p Packet) Clone() Packet {
|
func (p Packet) Clone() Packet {
|
||||||
if p.Bytes == nil {
|
if p.Bytes == nil {
|
||||||
return p
|
return p
|
||||||
}
|
}
|
||||||
cp := make([]byte, len(p.Bytes))
|
cp := make([]byte, len(p.Bytes))
|
||||||
copy(cp, p.Bytes)
|
copy(cp, p.Bytes)
|
||||||
return Packet{Bytes: cp}
|
return Packet{Bytes: cp, GSO: p.GSO}
|
||||||
|
}
|
||||||
|
|
||||||
|
// CapsProvider is an optional interface implemented by Queues that
|
||||||
|
// successfully negotiated kernel offload features at open time. Callers
|
||||||
|
// pick a write-path coalescer based on the result. Queues that don't
|
||||||
|
// implement it are treated as having no offload capability — callers must
|
||||||
|
// fall back to plain per-packet writes.
|
||||||
|
type CapsProvider interface {
|
||||||
|
Capabilities() Capabilities
|
||||||
|
}
|
||||||
|
|
||||||
|
// QueueCapabilities returns q's negotiated offload capabilities, or the
|
||||||
|
// zero value when q does not advertise any.
|
||||||
|
func QueueCapabilities(q Queue) Capabilities {
|
||||||
|
if cp, ok := q.(CapsProvider); ok {
|
||||||
|
return cp.Capabilities()
|
||||||
|
}
|
||||||
|
return Capabilities{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// GSOProto selects the L4 protocol for a GSO superpacket. Determines which
|
||||||
|
// VIRTIO_NET_HDR_GSO_* type the writer stamps and which checksum offset
|
||||||
|
// inside the transport header virtio NEEDS_CSUM expects.
|
||||||
|
type GSOProto uint8
|
||||||
|
|
||||||
|
const (
|
||||||
|
GSOProtoTCP GSOProto = iota
|
||||||
|
GSOProtoUDP
|
||||||
|
)
|
||||||
|
|
||||||
|
// GSOWriter is implemented by Queues that can emit a TCP or UDP superpacket
|
||||||
|
// assembled from a header prefix plus one or more borrowed payload
|
||||||
|
// fragments, in a single vectored write (writev with a leading
|
||||||
|
// virtio_net_hdr). This lets the coalescer avoid copying payload bytes
|
||||||
|
// between the caller's decrypt buffer and the TUN. Backends without GSO
|
||||||
|
// support do not implement this interface and coalescing is skipped.
|
||||||
|
//
|
||||||
|
// hdr contains the IPv4/IPv6 header prefix (mutable - callers will have
|
||||||
|
// filled in total length and IP csum). transportHdr is the TCP or UDP
|
||||||
|
// header (mutable - the L4 checksum field must hold the pseudo-header
|
||||||
|
// partial, single-fold not inverted, per virtio NEEDS_CSUM semantics).
|
||||||
|
// pays are non-overlapping payload fragments whose concatenation is the
|
||||||
|
// full superpacket payload; they are read-only from the writer's
|
||||||
|
// perspective and must remain valid until the call returns. Every segment
|
||||||
|
// in pays except possibly the last is exactly the same size. proto picks
|
||||||
|
// the L4 protocol so the writer knows which GSOType / CsumOffset to set.
|
||||||
|
//
|
||||||
|
// Callers should also consult CapsProvider (via SupportsGSO or
|
||||||
|
// QueueCapabilities) for the per-protocol negotiated capability; an
|
||||||
|
// implementation of GSOWriter is necessary but not sufficient since USO
|
||||||
|
// may not have been negotiated even when TSO was.
|
||||||
|
type GSOWriter interface {
|
||||||
|
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`. A writer that
|
||||||
|
// implements GSOWriter but not CapsProvider is treated as permissive
|
||||||
|
// (used by tests and fakes that don't negotiate).
|
||||||
|
func SupportsGSO(w any, want GSOProto) (GSOWriter, bool) {
|
||||||
|
gw, ok := w.(GSOWriter)
|
||||||
|
if !ok {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
cp, ok := w.(CapsProvider)
|
||||||
|
if !ok {
|
||||||
|
return gw, true
|
||||||
|
}
|
||||||
|
caps := cp.Capabilities()
|
||||||
|
switch want {
|
||||||
|
case GSOProtoTCP:
|
||||||
|
return gw, caps.TSO
|
||||||
|
case GSOProtoUDP:
|
||||||
|
return gw, caps.USO
|
||||||
|
}
|
||||||
|
return gw, false
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,414 @@
|
|||||||
|
//go:build linux && !android
|
||||||
|
// +build linux,!android
|
||||||
|
|
||||||
|
package tio
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"os"
|
||||||
|
"sync/atomic"
|
||||||
|
"syscall"
|
||||||
|
"unsafe"
|
||||||
|
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/overlay/tio/virtio"
|
||||||
|
)
|
||||||
|
|
||||||
|
// tunRxBufSize is the per-Read worst-case footprint inside rxBuf: one
|
||||||
|
// kernel-supplied packet body, which is at most ~64 KiB (tunReadBufSize).
|
||||||
|
// Segmentation happens at encrypt time on a per-routine MTU-sized scratch
|
||||||
|
// (see SegmentSuperpacket), so rxBuf only holds raw kernel-supplied bytes.
|
||||||
|
// We round up to give comfortable margin for the drain headroom check
|
||||||
|
// below.
|
||||||
|
const tunRxBufSize = 64 * 1024
|
||||||
|
|
||||||
|
// tunRxBufCap is the total size we allocate for the per-reader rx
|
||||||
|
// buffer. With reads landing directly in rxBuf, each drain iteration
|
||||||
|
// consumes up to tunRxBufSize of headroom for the kernel-supplied bytes.
|
||||||
|
// Sized to eight such iterations so a single poll wake can drain several
|
||||||
|
// TSO/USO superpackets under bulk load, amortizing the wake and giving
|
||||||
|
// the sendmmsg planner longer same-destination runs. Hold latency stays
|
||||||
|
// bounded because listenIn flushes its send batch incrementally rather
|
||||||
|
// than only at end-of-drain.
|
||||||
|
const tunRxBufCap = tunRxBufSize * 8
|
||||||
|
|
||||||
|
// tunDrainCap caps how many packets a single Read will accumulate via
|
||||||
|
// the post-wake drain loop. Sized to soak up a burst of small ACKs while
|
||||||
|
// bounding how much work a single caller holds before handing off.
|
||||||
|
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 — fine to keep resident on every queue. WriteGSO returns an error
|
||||||
|
// rather than reallocating when a caller exceeds this budget.
|
||||||
|
const gsoMaxIovs = 256
|
||||||
|
|
||||||
|
// validVnetHdr is the 10-byte virtio_net_hdr we prepend to every non-GSO TUN
|
||||||
|
// write. Only flag set is VIRTIO_NET_HDR_F_DATA_VALID, which marks the skb
|
||||||
|
// CHECKSUM_UNNECESSARY so the receiving network stack skips L4 checksum
|
||||||
|
// verification. All packets that reach the plain Write paths already carry
|
||||||
|
// a valid L4 checksum (either supplied by a remote peer whose ciphertext we
|
||||||
|
// AEAD-authenticated, produced by segmentTCPYield/segmentUDPYield during
|
||||||
|
// superpacket segmentation, or built locally by CreateRejectPacket), so
|
||||||
|
// trusting them is safe.
|
||||||
|
var validVnetHdr = [virtio.Size]byte{unix.VIRTIO_NET_HDR_F_DATA_VALID}
|
||||||
|
|
||||||
|
// Offload wraps a TUN file descriptor with poll-based reads. The FD provided will be changed to non-blocking.
|
||||||
|
// A shared eventfd allows Close to wake all readers blocked in poll.
|
||||||
|
type Offload struct {
|
||||||
|
fd int
|
||||||
|
shutdownFd int
|
||||||
|
closed atomic.Bool
|
||||||
|
rxBuf []byte // backing store for kernel-handed packets read this drain
|
||||||
|
rxOff int // cursor into rxBuf for the current Read drain
|
||||||
|
pending []Packet // packets returned from the most recent Read
|
||||||
|
|
||||||
|
// readVnetScratch holds the 10-byte virtio_net_hdr split off the front of
|
||||||
|
// every TUN read via readv(2). Decoupling the header from the packet body
|
||||||
|
// lets us read the body directly into rxBuf at the current rxOff with
|
||||||
|
// no userspace copy on the GSO_NONE fast path.
|
||||||
|
readVnetScratch [virtio.Size]byte
|
||||||
|
// readIovs is the readv(2) iovec scratch wired once at construction —
|
||||||
|
// iovec[0] points at readVnetScratch; iovec[1].Base/Len is updated per
|
||||||
|
// read to address the current rxBuf slot.
|
||||||
|
readIovs [2]unix.Iovec
|
||||||
|
|
||||||
|
// usoEnabled records whether the kernel agreed to TUN_F_USO* on this FD,
|
||||||
|
// so writers can decide whether emitting GSO_UDP_L4 superpackets is safe.
|
||||||
|
usoEnabled bool
|
||||||
|
|
||||||
|
// 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
|
||||||
|
}
|
||||||
|
|
||||||
|
func newOffload(fd int, shutdownFd int, usoEnabled bool) (*Offload, error) {
|
||||||
|
if err := unix.SetNonblock(fd, true); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
out := &Offload{
|
||||||
|
fd: fd,
|
||||||
|
shutdownFd: shutdownFd,
|
||||||
|
usoEnabled: usoEnabled,
|
||||||
|
closed: atomic.Bool{},
|
||||||
|
|
||||||
|
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.
|
||||||
|
//
|
||||||
|
// The body iovec capacity is always tunReadBufSize; callers (the Read
|
||||||
|
// drain loop) gate entry on tunRxBufCap-rxOff >= tunRxBufSize, sized to
|
||||||
|
// hold one worst-case kernel-supplied packet body. Without that gate the
|
||||||
|
// body iovec could be smaller than the next inbound packet and the
|
||||||
|
// kernel would truncate.
|
||||||
|
func (r *Offload) readPacket(block bool) (int, error) {
|
||||||
|
for {
|
||||||
|
r.readIovs[1].Base = &r.rxBuf[r.rxOff]
|
||||||
|
r.readIovs[1].SetLen(tunReadBufSize)
|
||||||
|
n, _, errno := syscall.Syscall(unix.SYS_READV, uintptr(r.fd), uintptr(unsafe.Pointer(&r.readIovs[0])), uintptr(len(r.readIovs)))
|
||||||
|
if errno == 0 {
|
||||||
|
if int(n) < virtio.Size {
|
||||||
|
return 0, io.ErrShortWrite
|
||||||
|
}
|
||||||
|
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.
|
||||||
|
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.
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return r.pending, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// decodeRead processes the packet sitting in rxBuf at rxOff (length
|
||||||
|
// pktLen). The bytes stay in rxBuf — for GSO_NONE we slice them as a
|
||||||
|
// regular IP datagram (running finishChecksum if NEEDS_CSUM is set);
|
||||||
|
// for TSO/USO superpackets we attach the corrected GSO metadata so the
|
||||||
|
// caller can segment lazily at encrypt time. rxOff advances past the
|
||||||
|
// kernel-supplied body and nothing else, since segmentation no longer
|
||||||
|
// writes back into rxBuf.
|
||||||
|
func (r *Offload) decodeRead(pktLen int) error {
|
||||||
|
if pktLen <= 0 {
|
||||||
|
return fmt.Errorf("short tun read: %d", pktLen)
|
||||||
|
}
|
||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
// GSO superpacket: validate, fix the kernel-supplied HdrLen on the
|
||||||
|
// FORWARD path (CorrectHdrLen), pick the L4 protocol, and attach
|
||||||
|
// the metadata. The bytes stay in rxBuf untouched, segmentation
|
||||||
|
// happens in SegmentSuperpacket at encrypt time.
|
||||||
|
if err := virtio.CheckValid(body, hdr); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := virtio.CorrectHdrLen(body, &hdr); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
proto, err := protoFromGSOType(hdr.GSOType)
|
||||||
|
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) {
|
||||||
|
iovs := [2]unix.Iovec{
|
||||||
|
{Base: &validVnetHdr[0]},
|
||||||
|
{Base: &buf[0]},
|
||||||
|
}
|
||||||
|
iovs[0].SetLen(virtio.Size)
|
||||||
|
iovs[1].SetLen(len(buf))
|
||||||
|
return r.writeWithScratch(buf, &iovs)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Offload) writeWithScratch(buf []byte, iovs *[2]unix.Iovec) (int, error) {
|
||||||
|
if len(buf) == 0 {
|
||||||
|
return 0, nil
|
||||||
|
}
|
||||||
|
iovs[1].Base = &buf[0]
|
||||||
|
iovs[1].SetLen(len(buf))
|
||||||
|
return r.rawWrite(unsafe.Slice(&iovs[0], len(iovs)))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
|
||||||
|
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(hdr) == 0 || len(pays) == 0 || len(transportHdr) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
// L4 checksum offset inside transportHdr: TCP=16 (the `check` field after
|
||||||
|
// seq/ack/dataoff/flags/window), UDP=6 (after sport/dport/length).
|
||||||
|
var csumOff uint16
|
||||||
|
switch proto {
|
||||||
|
case GSOProtoUDP:
|
||||||
|
csumOff = 6
|
||||||
|
default:
|
||||||
|
csumOff = 16
|
||||||
|
}
|
||||||
|
vhdr := virtio.Hdr{
|
||||||
|
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||||
|
HdrLen: uint16(len(hdr) + len(transportHdr)),
|
||||||
|
GSOSize: uint16(len(pays[0])),
|
||||||
|
CsumStart: uint16(len(hdr)),
|
||||||
|
CsumOffset: csumOff,
|
||||||
|
}
|
||||||
|
if len(pays) > 1 {
|
||||||
|
ipVer := hdr[0] >> 4
|
||||||
|
switch {
|
||||||
|
case proto == GSOProtoUDP && (ipVer == 4 || ipVer == 6):
|
||||||
|
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_UDP_L4
|
||||||
|
case ipVer == 6:
|
||||||
|
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV6
|
||||||
|
case ipVer == 4:
|
||||||
|
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV4
|
||||||
|
default:
|
||||||
|
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
||||||
|
vhdr.GSOSize = 0
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
||||||
|
vhdr.GSOSize = 0
|
||||||
|
}
|
||||||
|
vhdr.Encode(r.gsoHdrBuf[:])
|
||||||
|
|
||||||
|
// Build the iovec array: [virtio_hdr, hdr, transportHdr, pays...]. r.gsoIovs[0] is
|
||||||
|
// wired to gsoHdrBuf at construction and never changes.
|
||||||
|
need := 3 + len(pays)
|
||||||
|
if need > cap(r.gsoIovs) {
|
||||||
|
slog.Default().Warn("tio: WriteGSO iovec budget exceeded; dropping superpacket",
|
||||||
|
"need", need, "cap", cap(r.gsoIovs), "segments", len(pays))
|
||||||
|
return fmt.Errorf("tio: WriteGSO needs %d iovecs but cap is %d", need, cap(r.gsoIovs))
|
||||||
|
}
|
||||||
|
r.gsoIovs = r.gsoIovs[:need]
|
||||||
|
r.gsoIovs[1].Base = &hdr[0]
|
||||||
|
r.gsoIovs[1].SetLen(len(hdr))
|
||||||
|
r.gsoIovs[2].Base = &transportHdr[0]
|
||||||
|
r.gsoIovs[2].SetLen(len(transportHdr))
|
||||||
|
// Defense in depth: an empty payload fragment can't be a valid GSO
|
||||||
|
// segment and &p[0] would panic on it. Callers route zero-length
|
||||||
|
// datagrams through the plain path (see UDPCoalescer.commitParsed), so
|
||||||
|
// this should never fire, but skip empties rather than index into one.
|
||||||
|
// `n` tracks where the next payload iovec lands, since skips make it
|
||||||
|
// drift from 3+i.
|
||||||
|
n := 3
|
||||||
|
for _, p := range pays {
|
||||||
|
if len(p) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
r.gsoIovs[n].Base = &p[0]
|
||||||
|
r.gsoIovs[n].SetLen(len(p))
|
||||||
|
n++
|
||||||
|
}
|
||||||
|
r.gsoIovs = r.gsoIovs[:n]
|
||||||
|
|
||||||
|
_, 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 readOne, and mutating the field would race that load.
|
||||||
|
// It gets EBADF -> os.ErrClosed (or wakes via the shutdown eventfd's
|
||||||
|
// ppoll first). closed.Swap already guarantees we only close once.
|
||||||
|
return unix.Close(r.fd)
|
||||||
|
}
|
||||||
@@ -11,8 +11,9 @@ import (
|
|||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Maximum size we accept for a single read from a TUN. 65535 covers any
|
// Maximum size we accept for a single read from a TUN with IFF_VNET_HDR. A
|
||||||
// single IP packet.
|
// TSO superpacket can be up to 64KiB of payload plus a single L2/L3/L4 header
|
||||||
|
// prefix plus the virtio header.
|
||||||
const tunReadBufSize = 65535
|
const tunReadBufSize = 65535
|
||||||
|
|
||||||
type Poll struct {
|
type Poll struct {
|
||||||
@@ -26,8 +27,8 @@ type Poll struct {
|
|||||||
|
|
||||||
// newPoll wraps an existing tun fd. On failure it does NOT close fd: the
|
// 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
|
// caller owns fd and is the sole closer (see pollQueueSet.Add callers in
|
||||||
// overlay/tun_linux.go, which unix.Close on Add error). This keeps closes
|
// overlay/tun_linux.go, which unix.Close on Add error). This matches the
|
||||||
// at exactly one on every path.
|
// newOffload convention and keeps closes at exactly one on every path.
|
||||||
func newPoll(fd int, shutdownFd int) (*Poll, error) {
|
func newPoll(fd int, shutdownFd int) (*Poll, error) {
|
||||||
if err := unix.SetNonblock(fd, true); err != nil {
|
if err := unix.SetNonblock(fd, true); err != nil {
|
||||||
return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err)
|
return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err)
|
||||||
|
|||||||
@@ -206,3 +206,22 @@ func TestPollQueueSet_Close_ClosesShutdownFd(t *testing.T) {
|
|||||||
// Second Close must not touch fds (shutdownFd is now -1) and must return nil.
|
// Second Close must not touch fds (shutdownFd is now -1) and must return nil.
|
||||||
require.NoError(t, qs.Close())
|
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)
|
||||||
|
require.NoError(t, err)
|
||||||
|
c, ok := qs.(*offloadQueueSet)
|
||||||
|
require.True(t, ok)
|
||||||
|
require.NoError(t, qs.Add(newReadPipe(t)))
|
||||||
|
|
||||||
|
shutdownFd := c.shutdownFd
|
||||||
|
require.True(t, fdOpen(t, shutdownFd), "shutdown eventfd should be open before Close")
|
||||||
|
|
||||||
|
require.NoError(t, qs.Close())
|
||||||
|
require.False(t, fdOpen(t, shutdownFd), "shutdown eventfd should be closed after Close")
|
||||||
|
|
||||||
|
// Second Close must not touch fds (shutdownFd is now -1) and must return nil.
|
||||||
|
require.NoError(t, qs.Close())
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,51 @@
|
|||||||
|
//go:build linux && !android && !e2e_testing
|
||||||
|
// +build linux,!android,!e2e_testing
|
||||||
|
|
||||||
|
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 {
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SegmentSuperpacket invokes fn once per segment of pkt. For non-GSO pkts
|
||||||
|
// fn is called once with pkt.Bytes (no segmentation, no copy). For GSO/USO
|
||||||
|
// superpackets fn is called once per segment with a slice of pkt.Bytes
|
||||||
|
// holding that segment's plaintext (a freshly-patched L3+L4 header sliced
|
||||||
|
// in front of the original payload chunk). The slide is destructive: pkt is
|
||||||
|
// consumed by this call and its bytes are in an undefined state when
|
||||||
|
// SegmentSuperpacket returns. Callers must not retain pkt or any earlier
|
||||||
|
// seg slice past fn's return for that segment. The scratch parameter is
|
||||||
|
// unused on the destructive path and kept only for cross-platform
|
||||||
|
// signature compatibility. Aborts and returns the first error from fn or
|
||||||
|
// from per-segment construction.
|
||||||
|
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
||||||
|
if !pkt.GSO.IsSuperpacket() {
|
||||||
|
return fn(pkt.Bytes)
|
||||||
|
}
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,826 @@
|
|||||||
|
//go:build linux && !android && !e2e_testing
|
||||||
|
// +build linux,!android,!e2e_testing
|
||||||
|
|
||||||
|
package tio
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
"gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/overlay/tio/virtio"
|
||||||
|
)
|
||||||
|
|
||||||
|
// testSegScratchSize is a generous segmentation scratch sized to fit any
|
||||||
|
// of the synthetic TSO/USO superpackets these tests generate (one
|
||||||
|
// worst-case 64 KiB superpacket plus replicated per-segment headers).
|
||||||
|
const testSegScratchSize = 192 * 1024
|
||||||
|
|
||||||
|
// verifyChecksum confirms that the one's-complement sum across `b`, seeded
|
||||||
|
// with a folded pseudo-header sum, equals all-ones (valid).
|
||||||
|
func verifyChecksum(b []byte, pseudo uint16) bool {
|
||||||
|
return checksum.Checksum(b, pseudo) == 0xffff
|
||||||
|
}
|
||||||
|
|
||||||
|
// segmentForTest is the test-only counterpart to the production
|
||||||
|
// SegmentSuperpacket path. It handles GSO_NONE (with optional
|
||||||
|
// finishChecksum) inline and dispatches GSO superpackets through
|
||||||
|
// SegmentSuperpacket, draining each yielded segment into a
|
||||||
|
// freshly-copied [][]byte slot so callers can iterate after the call
|
||||||
|
// returns. Tests pre-set hdr.HdrLen correctly, so correctHdrLen is not
|
||||||
|
// invoked here.
|
||||||
|
func segmentForTest(pkt []byte, hdr virtio.Hdr, out *[][]byte, scratch []byte) error {
|
||||||
|
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||||
|
cp := append([]byte(nil), pkt...)
|
||||||
|
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
|
||||||
|
if err := virtio.FinishChecksum(cp, hdr); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
*out = append(*out, cp)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
proto, err := protoFromGSOType(hdr.GSOType)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
gso := GSOInfo{
|
||||||
|
Size: hdr.GSOSize,
|
||||||
|
HdrLen: hdr.HdrLen,
|
||||||
|
CsumStart: hdr.CsumStart,
|
||||||
|
Proto: proto,
|
||||||
|
}
|
||||||
|
return SegmentSuperpacket(Packet{Bytes: pkt, GSO: gso}, func(seg []byte) error {
|
||||||
|
*out = append(*out, append([]byte(nil), seg...))
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// pseudoHeaderIPv4 returns the folded pseudo-header sum used to verify a
|
||||||
|
// TCP/UDP segment's checksum in tests. src/dst are 4 bytes each.
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
|
||||||
|
// pseudoHeaderIPv6 returns the folded pseudo-header sum used to verify a
|
||||||
|
// TCP/UDP segment's checksum in tests. src/dst are 16 bytes each.
|
||||||
|
func pseudoHeaderIPv6(src, dst []byte, proto byte, l4Len int) uint16 {
|
||||||
|
s := uint32(checksum.Checksum(src, 0)) + uint32(checksum.Checksum(dst, 0))
|
||||||
|
s += uint32(l4Len>>16) + uint32(l4Len&0xffff) + uint32(proto)
|
||||||
|
s = (s & 0xffff) + (s >> 16)
|
||||||
|
s = (s & 0xffff) + (s >> 16)
|
||||||
|
return uint16(s)
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildTSOv4 builds a synthetic IPv4/TCP TSO superpacket with a payload of
|
||||||
|
// `payLen` bytes split at `mss`.
|
||||||
|
func buildTSOv4(t *testing.T, payLen, mss int) ([]byte, virtio.Hdr) {
|
||||||
|
t.Helper()
|
||||||
|
const ipLen = 20
|
||||||
|
const tcpLen = 20
|
||||||
|
pkt := make([]byte, ipLen+tcpLen+payLen)
|
||||||
|
|
||||||
|
// IPv4 header
|
||||||
|
pkt[0] = 0x45 // version 4, IHL 5
|
||||||
|
// total length is meaningless for TSO but set it anyway
|
||||||
|
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+tcpLen+payLen))
|
||||||
|
binary.BigEndian.PutUint16(pkt[4:6], 0x4242) // original 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
|
||||||
|
|
||||||
|
// payload
|
||||||
|
for i := 0; i < payLen; i++ {
|
||||||
|
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
|
||||||
|
}
|
||||||
|
|
||||||
|
return pkt, virtio.Hdr{
|
||||||
|
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||||
|
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4,
|
||||||
|
HdrLen: uint16(ipLen + tcpLen),
|
||||||
|
GSOSize: uint16(mss),
|
||||||
|
CsumStart: uint16(ipLen),
|
||||||
|
CsumOffset: 16,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSegmentTCPv4(t *testing.T) {
|
||||||
|
const mss = 100
|
||||||
|
const numSeg = 3
|
||||||
|
pkt, hdr := buildTSOv4(t, mss*numSeg, mss)
|
||||||
|
|
||||||
|
scratch := make([]byte, testSegScratchSize)
|
||||||
|
var out [][]byte
|
||||||
|
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
||||||
|
t.Fatalf("segmentForTest: %v", err)
|
||||||
|
}
|
||||||
|
if len(out) != numSeg {
|
||||||
|
t.Fatalf("expected %d segments, got %d", numSeg, len(out))
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, seg := range out {
|
||||||
|
if len(seg) != 40+mss {
|
||||||
|
t.Errorf("seg %d: unexpected len %d", i, len(seg))
|
||||||
|
}
|
||||||
|
totalLen := binary.BigEndian.Uint16(seg[2:4])
|
||||||
|
if totalLen != uint16(40+mss) {
|
||||||
|
t.Errorf("seg %d: total_len=%d want %d", i, totalLen, 40+mss)
|
||||||
|
}
|
||||||
|
id := binary.BigEndian.Uint16(seg[4:6])
|
||||||
|
if id != 0x4242+uint16(i) {
|
||||||
|
t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242+uint16(i))
|
||||||
|
}
|
||||||
|
seq := binary.BigEndian.Uint32(seg[24:28])
|
||||||
|
wantSeq := uint32(10000 + i*mss)
|
||||||
|
if seq != wantSeq {
|
||||||
|
t.Errorf("seg %d: seq=%d want %d", i, seq, wantSeq)
|
||||||
|
}
|
||||||
|
flags := seg[33]
|
||||||
|
wantFlags := byte(0x10) // ACK only, PSH cleared
|
||||||
|
if i == numSeg-1 {
|
||||||
|
wantFlags = 0x18 // ACK | PSH preserved on last
|
||||||
|
}
|
||||||
|
if flags != wantFlags {
|
||||||
|
t.Errorf("seg %d: flags=%#x want %#x", i, flags, wantFlags)
|
||||||
|
}
|
||||||
|
// IPv4 header checksum must verify against itself.
|
||||||
|
if !verifyChecksum(seg[:20], 0) {
|
||||||
|
t.Errorf("seg %d: bad IPv4 header checksum", i)
|
||||||
|
}
|
||||||
|
// TCP checksum must verify against the pseudo-header.
|
||||||
|
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_TCP, 20+mss)
|
||||||
|
if !verifyChecksum(seg[20:], psum) {
|
||||||
|
t.Errorf("seg %d: bad TCP checksum", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSegmentTCPv4OddTail(t *testing.T) {
|
||||||
|
// Payload of 250 bytes with MSS 100 → segments of 100, 100, 50.
|
||||||
|
pkt, hdr := buildTSOv4(t, 250, 100)
|
||||||
|
scratch := make([]byte, testSegScratchSize)
|
||||||
|
var out [][]byte
|
||||||
|
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
||||||
|
t.Fatalf("segmentForTest: %v", err)
|
||||||
|
}
|
||||||
|
if len(out) != 3 {
|
||||||
|
t.Fatalf("want 3 segments, got %d", len(out))
|
||||||
|
}
|
||||||
|
wantPayLens := []int{100, 100, 50}
|
||||||
|
for i, seg := range out {
|
||||||
|
if len(seg)-40 != wantPayLens[i] {
|
||||||
|
t.Errorf("seg %d: pay len %d want %d", i, len(seg)-40, wantPayLens[i])
|
||||||
|
}
|
||||||
|
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, 20+wantPayLens[i])
|
||||||
|
if !verifyChecksum(seg[20:], psum) {
|
||||||
|
t.Errorf("seg %d: bad TCP checksum", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSegmentTCPv6(t *testing.T) {
|
||||||
|
const ipLen = 40
|
||||||
|
const tcpLen = 20
|
||||||
|
const mss = 120
|
||||||
|
const numSeg = 2
|
||||||
|
payLen := mss * numSeg
|
||||||
|
pkt := make([]byte, ipLen+tcpLen+payLen)
|
||||||
|
|
||||||
|
// IPv6 header
|
||||||
|
pkt[0] = 0x60 // version 6
|
||||||
|
binary.BigEndian.PutUint16(pkt[4:6], uint16(tcpLen+payLen))
|
||||||
|
pkt[6] = unix.IPPROTO_TCP
|
||||||
|
pkt[7] = 64
|
||||||
|
// src/dst fe80::1 / fe80::2
|
||||||
|
pkt[8] = 0xfe
|
||||||
|
pkt[9] = 0x80
|
||||||
|
pkt[23] = 1
|
||||||
|
pkt[24] = 0xfe
|
||||||
|
pkt[25] = 0x80
|
||||||
|
pkt[39] = 2
|
||||||
|
|
||||||
|
// TCP header
|
||||||
|
binary.BigEndian.PutUint16(pkt[40:42], 12345)
|
||||||
|
binary.BigEndian.PutUint16(pkt[42:44], 80)
|
||||||
|
binary.BigEndian.PutUint32(pkt[44:48], 7)
|
||||||
|
binary.BigEndian.PutUint32(pkt[48:52], 99)
|
||||||
|
pkt[52] = 0x50
|
||||||
|
pkt[53] = 0x19 // FIN | ACK | PSH — exercise FIN clearing too
|
||||||
|
binary.BigEndian.PutUint16(pkt[54:56], 65535)
|
||||||
|
|
||||||
|
for i := 0; i < payLen; i++ {
|
||||||
|
pkt[ipLen+tcpLen+i] = byte(i)
|
||||||
|
}
|
||||||
|
|
||||||
|
hdr := virtio.Hdr{
|
||||||
|
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||||
|
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV6,
|
||||||
|
HdrLen: uint16(ipLen + tcpLen),
|
||||||
|
GSOSize: uint16(mss),
|
||||||
|
CsumStart: uint16(ipLen),
|
||||||
|
CsumOffset: 16,
|
||||||
|
}
|
||||||
|
|
||||||
|
scratch := make([]byte, testSegScratchSize)
|
||||||
|
var out [][]byte
|
||||||
|
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
||||||
|
t.Fatalf("segmentForTest: %v", err)
|
||||||
|
}
|
||||||
|
if len(out) != numSeg {
|
||||||
|
t.Fatalf("want %d segments, got %d", numSeg, len(out))
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, seg := range out {
|
||||||
|
if len(seg) != ipLen+tcpLen+mss {
|
||||||
|
t.Errorf("seg %d: len %d want %d", i, len(seg), ipLen+tcpLen+mss)
|
||||||
|
}
|
||||||
|
pl := binary.BigEndian.Uint16(seg[4:6])
|
||||||
|
if pl != uint16(tcpLen+mss) {
|
||||||
|
t.Errorf("seg %d: payload_length=%d want %d", i, pl, tcpLen+mss)
|
||||||
|
}
|
||||||
|
seq := binary.BigEndian.Uint32(seg[44:48])
|
||||||
|
if seq != uint32(7+i*mss) {
|
||||||
|
t.Errorf("seg %d: seq=%d want %d", i, seq, 7+i*mss)
|
||||||
|
}
|
||||||
|
flags := seg[53]
|
||||||
|
// Original flags = 0x19 (FIN|ACK|PSH). FIN(0x01)+PSH(0x08) should be
|
||||||
|
// cleared on all but the last; ACK(0x10) always preserved.
|
||||||
|
wantFlags := byte(0x10)
|
||||||
|
if i == numSeg-1 {
|
||||||
|
wantFlags = 0x19
|
||||||
|
}
|
||||||
|
if flags != wantFlags {
|
||||||
|
t.Errorf("seg %d: flags=%#x want %#x", i, flags, wantFlags)
|
||||||
|
}
|
||||||
|
psum := pseudoHeaderIPv6(seg[8:24], seg[24:40], unix.IPPROTO_TCP, tcpLen+mss)
|
||||||
|
if !verifyChecksum(seg[ipLen:], psum) {
|
||||||
|
t.Errorf("seg %d: bad TCP checksum", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSegmentGSONonePassesThrough(t *testing.T) {
|
||||||
|
pkt, hdr := buildTSOv4(t, 100, 100)
|
||||||
|
hdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
||||||
|
hdr.Flags = 0 // no NEEDS_CSUM, leave packet untouched
|
||||||
|
|
||||||
|
scratch := make([]byte, testSegScratchSize)
|
||||||
|
var out [][]byte
|
||||||
|
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
||||||
|
t.Fatalf("segmentForTest: %v", err)
|
||||||
|
}
|
||||||
|
if len(out) != 1 {
|
||||||
|
t.Fatalf("want 1 segment, got %d", len(out))
|
||||||
|
}
|
||||||
|
if len(out[0]) != len(pkt) {
|
||||||
|
t.Fatalf("unexpected length: %d vs %d", len(out[0]), len(pkt))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSegmentRejectsLegacyUDPGSO ensures the legacy GSO_UDP (UFO) marker is
|
||||||
|
// still rejected; only modern GSO_UDP_L4 (USO) is supported.
|
||||||
|
func TestSegmentRejectsLegacyUDPGSO(t *testing.T) {
|
||||||
|
hdr := virtio.Hdr{GSOType: unix.VIRTIO_NET_HDR_GSO_UDP}
|
||||||
|
var out [][]byte
|
||||||
|
if err := segmentForTest(nil, hdr, &out, nil); err == nil {
|
||||||
|
t.Fatalf("expected rejection for legacy UDP GSO")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildUSOv4 builds a synthetic IPv4/UDP USO superpacket with payload of
|
||||||
|
// payLen bytes, segmented at gsoSize.
|
||||||
|
func buildUSOv4(t *testing.T, payLen, gsoSize int) ([]byte, virtio.Hdr) {
|
||||||
|
t.Helper()
|
||||||
|
const ipLen = 20
|
||||||
|
const udpLen = 8
|
||||||
|
pkt := make([]byte, ipLen+udpLen+payLen)
|
||||||
|
|
||||||
|
// IPv4 header
|
||||||
|
pkt[0] = 0x45 // version 4, IHL 5
|
||||||
|
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})
|
||||||
|
|
||||||
|
// UDP header (length + checksum filled in per segment by segmentUDPYield)
|
||||||
|
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
|
||||||
|
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
|
||||||
|
|
||||||
|
for i := 0; i < payLen; i++ {
|
||||||
|
pkt[ipLen+udpLen+i] = byte(i & 0xff)
|
||||||
|
}
|
||||||
|
|
||||||
|
return pkt, virtio.Hdr{
|
||||||
|
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||||
|
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
|
||||||
|
HdrLen: uint16(ipLen + udpLen),
|
||||||
|
GSOSize: uint16(gsoSize),
|
||||||
|
CsumStart: uint16(ipLen),
|
||||||
|
CsumOffset: 6,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSegmentUDPv4(t *testing.T) {
|
||||||
|
const gso = 100
|
||||||
|
const numSeg = 3
|
||||||
|
pkt, hdr := buildUSOv4(t, gso*numSeg, gso)
|
||||||
|
|
||||||
|
scratch := make([]byte, testSegScratchSize)
|
||||||
|
var out [][]byte
|
||||||
|
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
||||||
|
t.Fatalf("segmentForTest: %v", err)
|
||||||
|
}
|
||||||
|
if len(out) != numSeg {
|
||||||
|
t.Fatalf("expected %d segments, got %d", numSeg, len(out))
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, seg := range out {
|
||||||
|
if len(seg) != 28+gso {
|
||||||
|
t.Errorf("seg %d: len %d want %d", i, len(seg), 28+gso)
|
||||||
|
}
|
||||||
|
totalLen := binary.BigEndian.Uint16(seg[2:4])
|
||||||
|
if totalLen != uint16(28+gso) {
|
||||||
|
t.Errorf("seg %d: total_len=%d want %d", i, totalLen, 28+gso)
|
||||||
|
}
|
||||||
|
// kernel UDP-GSO does NOT bump the IPv4 ID across segments; every
|
||||||
|
// segment carries the same ID as the seed.
|
||||||
|
id := binary.BigEndian.Uint16(seg[4:6])
|
||||||
|
if id != 0x4242 {
|
||||||
|
t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242)
|
||||||
|
}
|
||||||
|
udpLen := binary.BigEndian.Uint16(seg[24:26])
|
||||||
|
if udpLen != uint16(8+gso) {
|
||||||
|
t.Errorf("seg %d: udp len=%d want %d", i, udpLen, 8+gso)
|
||||||
|
}
|
||||||
|
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, 8+gso)
|
||||||
|
if !verifyChecksum(seg[20:], psum) {
|
||||||
|
t.Errorf("seg %d: bad UDP checksum", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSegmentUDPv4OddTail(t *testing.T) {
|
||||||
|
// 250 bytes payload, gsoSize=100 → segments of 100, 100, 50.
|
||||||
|
pkt, hdr := buildUSOv4(t, 250, 100)
|
||||||
|
scratch := make([]byte, testSegScratchSize)
|
||||||
|
var out [][]byte
|
||||||
|
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
||||||
|
t.Fatalf("segmentForTest: %v", err)
|
||||||
|
}
|
||||||
|
if len(out) != 3 {
|
||||||
|
t.Fatalf("want 3 segments, got %d", len(out))
|
||||||
|
}
|
||||||
|
wantPay := []int{100, 100, 50}
|
||||||
|
for i, seg := range out {
|
||||||
|
if len(seg)-28 != wantPay[i] {
|
||||||
|
t.Errorf("seg %d: pay len %d want %d", i, len(seg)-28, wantPay[i])
|
||||||
|
}
|
||||||
|
udpLen := binary.BigEndian.Uint16(seg[24:26])
|
||||||
|
if udpLen != uint16(8+wantPay[i]) {
|
||||||
|
t.Errorf("seg %d: udp len=%d want %d", i, udpLen, 8+wantPay[i])
|
||||||
|
}
|
||||||
|
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, 8+wantPay[i])
|
||||||
|
if !verifyChecksum(seg[20:], psum) {
|
||||||
|
t.Errorf("seg %d: bad UDP checksum", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSegmentUDPv6(t *testing.T) {
|
||||||
|
const ipLen = 40
|
||||||
|
const udpLen = 8
|
||||||
|
const gso = 120
|
||||||
|
const numSeg = 2
|
||||||
|
payLen := gso * numSeg
|
||||||
|
pkt := make([]byte, ipLen+udpLen+payLen)
|
||||||
|
|
||||||
|
// IPv6 header
|
||||||
|
pkt[0] = 0x60
|
||||||
|
binary.BigEndian.PutUint16(pkt[4:6], uint16(udpLen+payLen))
|
||||||
|
pkt[6] = unix.IPPROTO_UDP
|
||||||
|
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], 12345)
|
||||||
|
binary.BigEndian.PutUint16(pkt[42:44], 53)
|
||||||
|
|
||||||
|
for i := 0; i < payLen; i++ {
|
||||||
|
pkt[ipLen+udpLen+i] = byte(i)
|
||||||
|
}
|
||||||
|
|
||||||
|
hdr := virtio.Hdr{
|
||||||
|
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||||
|
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
|
||||||
|
HdrLen: uint16(ipLen + udpLen),
|
||||||
|
GSOSize: uint16(gso),
|
||||||
|
CsumStart: uint16(ipLen),
|
||||||
|
CsumOffset: 6,
|
||||||
|
}
|
||||||
|
|
||||||
|
scratch := make([]byte, testSegScratchSize)
|
||||||
|
var out [][]byte
|
||||||
|
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
||||||
|
t.Fatalf("segmentForTest: %v", err)
|
||||||
|
}
|
||||||
|
if len(out) != numSeg {
|
||||||
|
t.Fatalf("want %d segments, got %d", numSeg, len(out))
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, seg := range out {
|
||||||
|
if len(seg) != ipLen+udpLen+gso {
|
||||||
|
t.Errorf("seg %d: len %d want %d", i, len(seg), ipLen+udpLen+gso)
|
||||||
|
}
|
||||||
|
pl := binary.BigEndian.Uint16(seg[4:6])
|
||||||
|
if pl != uint16(udpLen+gso) {
|
||||||
|
t.Errorf("seg %d: payload_length=%d want %d", i, pl, udpLen+gso)
|
||||||
|
}
|
||||||
|
ul := binary.BigEndian.Uint16(seg[ipLen+4 : ipLen+6])
|
||||||
|
if ul != uint16(udpLen+gso) {
|
||||||
|
t.Errorf("seg %d: udp len=%d want %d", i, ul, udpLen+gso)
|
||||||
|
}
|
||||||
|
psum := pseudoHeaderIPv6(seg[8:24], seg[24:40], unix.IPPROTO_UDP, udpLen+gso)
|
||||||
|
if !verifyChecksum(seg[ipLen:], psum) {
|
||||||
|
t.Errorf("seg %d: bad UDP checksum", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSegmentUDPCEPropagates confirms IP-level CE marks on the seed appear on
|
||||||
|
// every segment. UDP has no transport-level CWR/ECE: the IP TOS/TC byte is
|
||||||
|
// copied verbatim into every segment by the segment-prefix copy.
|
||||||
|
func TestSegmentUDPCEPropagates(t *testing.T) {
|
||||||
|
pkt, hdr := buildUSOv4(t, 200, 100)
|
||||||
|
pkt[1] = 0x03 // CE codepoint in IP-ECN
|
||||||
|
|
||||||
|
scratch := make([]byte, testSegScratchSize)
|
||||||
|
var out [][]byte
|
||||||
|
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
||||||
|
t.Fatalf("segmentForTest: %v", err)
|
||||||
|
}
|
||||||
|
if len(out) != 2 {
|
||||||
|
t.Fatalf("want 2 segments, got %d", len(out))
|
||||||
|
}
|
||||||
|
for i, seg := range out {
|
||||||
|
if seg[1]&0x03 != 0x03 {
|
||||||
|
t.Errorf("seg %d: CE missing (tos=%#x)", i, seg[1])
|
||||||
|
}
|
||||||
|
if !verifyChecksum(seg[:20], 0) {
|
||||||
|
t.Errorf("seg %d: bad IPv4 header checksum", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSegmentTCPCwrFirstSegmentOnly confirms RFC 3168 §6.1.2: when a TSO
|
||||||
|
// burst's seed has CWR set, only the first emitted segment carries CWR.
|
||||||
|
// ECE is preserved on every segment (different signal, persistent state).
|
||||||
|
func TestSegmentTCPCwrFirstSegmentOnly(t *testing.T) {
|
||||||
|
const mss = 100
|
||||||
|
const numSeg = 3
|
||||||
|
pkt, hdr := buildTSOv4(t, mss*numSeg, mss)
|
||||||
|
// Seed flags: CWR | ECE | ACK | PSH.
|
||||||
|
pkt[33] = 0x80 | 0x40 | 0x10 | 0x08
|
||||||
|
|
||||||
|
scratch := make([]byte, testSegScratchSize)
|
||||||
|
var out [][]byte
|
||||||
|
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
||||||
|
t.Fatalf("segmentForTest: %v", err)
|
||||||
|
}
|
||||||
|
if len(out) != numSeg {
|
||||||
|
t.Fatalf("expected %d segments, got %d", numSeg, len(out))
|
||||||
|
}
|
||||||
|
for i, seg := range out {
|
||||||
|
flags := seg[33]
|
||||||
|
hasCwr := flags&0x80 != 0
|
||||||
|
hasEce := flags&0x40 != 0
|
||||||
|
hasPsh := flags&0x08 != 0
|
||||||
|
wantCwr := i == 0
|
||||||
|
wantPsh := i == numSeg-1
|
||||||
|
if hasCwr != wantCwr {
|
||||||
|
t.Errorf("seg %d: CWR=%v want %v (flags=%#x)", i, hasCwr, wantCwr, flags)
|
||||||
|
}
|
||||||
|
if !hasEce {
|
||||||
|
t.Errorf("seg %d: ECE missing (flags=%#x)", i, flags)
|
||||||
|
}
|
||||||
|
if hasPsh != wantPsh {
|
||||||
|
t.Errorf("seg %d: PSH=%v want %v (flags=%#x)", i, hasPsh, wantPsh, flags)
|
||||||
|
}
|
||||||
|
// IP and TCP checksums must still verify after the flag rewrite.
|
||||||
|
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, 20+mss)
|
||||||
|
if !verifyChecksum(seg[20:], psum) {
|
||||||
|
t.Errorf("seg %d: bad TCP checksum", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func BenchmarkSegmentTCPv4(b *testing.B) {
|
||||||
|
sizes := []struct {
|
||||||
|
name string
|
||||||
|
payLen int
|
||||||
|
mss int
|
||||||
|
}{
|
||||||
|
{"64KiB_MSS1460", 65000, 1460},
|
||||||
|
{"16KiB_MSS1460", 16384, 1460},
|
||||||
|
{"4KiB_MSS1460", 4096, 1460},
|
||||||
|
}
|
||||||
|
for _, sz := range sizes {
|
||||||
|
b.Run(sz.name, func(b *testing.B) {
|
||||||
|
const ipLen = 20
|
||||||
|
const tcpLen = 20
|
||||||
|
pkt := make([]byte, ipLen+tcpLen+sz.payLen)
|
||||||
|
pkt[0] = 0x45
|
||||||
|
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+tcpLen+sz.payLen))
|
||||||
|
binary.BigEndian.PutUint16(pkt[4:6], 0x4242)
|
||||||
|
pkt[8] = 64
|
||||||
|
pkt[9] = unix.IPPROTO_TCP
|
||||||
|
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||||
|
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||||
|
binary.BigEndian.PutUint16(pkt[20:22], 12345)
|
||||||
|
binary.BigEndian.PutUint16(pkt[22:24], 80)
|
||||||
|
binary.BigEndian.PutUint32(pkt[24:28], 10000)
|
||||||
|
binary.BigEndian.PutUint32(pkt[28:32], 20000)
|
||||||
|
pkt[32] = 0x50
|
||||||
|
pkt[33] = 0x18
|
||||||
|
binary.BigEndian.PutUint16(pkt[34:36], 65535)
|
||||||
|
for i := 0; i < sz.payLen; i++ {
|
||||||
|
pkt[ipLen+tcpLen+i] = byte(i)
|
||||||
|
}
|
||||||
|
hdr := virtio.Hdr{
|
||||||
|
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||||
|
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4,
|
||||||
|
HdrLen: uint16(ipLen + tcpLen),
|
||||||
|
GSOSize: uint16(sz.mss),
|
||||||
|
CsumStart: uint16(ipLen),
|
||||||
|
CsumOffset: 16,
|
||||||
|
}
|
||||||
|
|
||||||
|
scratch := make([]byte, testSegScratchSize)
|
||||||
|
out := make([][]byte, 0, 64)
|
||||||
|
|
||||||
|
// SegmentSuperpacket consumes its input destructively; restore
|
||||||
|
// pkt from a master copy each iteration. The restore mirrors the
|
||||||
|
// kernel→userspace copy that hands a fresh GSO blob to the
|
||||||
|
// segmenter in production, so it's representative cost rather
|
||||||
|
// than bench overhead.
|
||||||
|
master := append([]byte(nil), pkt...)
|
||||||
|
work := make([]byte, len(pkt))
|
||||||
|
|
||||||
|
b.SetBytes(int64(len(pkt)))
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
copy(work, master)
|
||||||
|
out = out[:0]
|
||||||
|
if err := segmentForTest(work, hdr, &out, scratch); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestTunFileWriteVnetHdrNoAlloc verifies the IFF_VNET_HDR fast-path write is
|
||||||
|
// allocation-free. We write to /dev/null so every call succeeds synchronously.
|
||||||
|
func TestTunFileWriteVnetHdrNoAlloc(t *testing.T) {
|
||||||
|
fd, err := unix.Open("/dev/null", os.O_WRONLY, 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("open /dev/null: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { _ = unix.Close(fd) })
|
||||||
|
|
||||||
|
tf := &Offload{fd: fd}
|
||||||
|
|
||||||
|
payload := make([]byte, 1400)
|
||||||
|
// Warm up (first call may trigger one-time internal allocations elsewhere).
|
||||||
|
if _, err := tf.Write(payload); err != nil {
|
||||||
|
t.Fatalf("Write: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
allocs := testing.AllocsPerRun(1000, func() {
|
||||||
|
if _, err := tf.Write(payload); err != nil {
|
||||||
|
t.Fatalf("Write: %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
if allocs != 0 {
|
||||||
|
t.Fatalf("Write allocated %.1f times per call, want 0", allocs)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWriteGSOSkipsEmptyPayloads is the defense-in-depth guard for the
|
||||||
|
// zero-length UDP DoS: a payload fragment of length zero would make &p[0]
|
||||||
|
// panic (index-out-of-range) when building the iovec array. WriteGSO must
|
||||||
|
// skip empties instead. We write to /dev/null so the writev always succeeds
|
||||||
|
// synchronously; the point is simply that neither call panics.
|
||||||
|
func TestWriteGSOSkipsEmptyPayloads(t *testing.T) {
|
||||||
|
fd, err := unix.Open("/dev/null", os.O_WRONLY, 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("open /dev/null: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { _ = unix.Close(fd) })
|
||||||
|
|
||||||
|
o := &Offload{fd: fd, gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs)}
|
||||||
|
o.gsoIovs[0].Base = &o.gsoHdrBuf[0]
|
||||||
|
o.gsoIovs[0].SetLen(virtio.Size)
|
||||||
|
|
||||||
|
ipHdr := make([]byte, 20)
|
||||||
|
ipHdr[0] = 0x45 // IPv4, IHL 5
|
||||||
|
udpHdr := make([]byte, 8)
|
||||||
|
|
||||||
|
// Sole payload empty: exercises the all-empty skip (n stays at 3).
|
||||||
|
if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{{}}, GSOProtoUDP); err != nil {
|
||||||
|
t.Fatalf("WriteGSO with a single empty payload: %v", err)
|
||||||
|
}
|
||||||
|
// Empty mixed with a real fragment: exercises the index-drift skip so a
|
||||||
|
// later non-empty payload still lands in the right iovec slot.
|
||||||
|
real := make([]byte, 1200)
|
||||||
|
if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{real, {}}, GSOProtoUDP); err != nil {
|
||||||
|
t.Fatalf("WriteGSO with a trailing empty payload: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildTSOv6 builds a synthetic IPv6/TCP TSO superpacket with payLen bytes
|
||||||
|
// of payload, segmented at gso. Returns the packet bytes only; the
|
||||||
|
// virtio_net_hdr is the caller's responsibility.
|
||||||
|
func buildTSOv6(payLen, gso int) []byte {
|
||||||
|
const ipLen = 40
|
||||||
|
const tcpLen = 20
|
||||||
|
pkt := make([]byte, ipLen+tcpLen+payLen)
|
||||||
|
|
||||||
|
pkt[0] = 0x60 // version 6
|
||||||
|
binary.BigEndian.PutUint16(pkt[4:6], uint16(tcpLen+payLen))
|
||||||
|
pkt[6] = unix.IPPROTO_TCP
|
||||||
|
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], 12345)
|
||||||
|
binary.BigEndian.PutUint16(pkt[42:44], 80)
|
||||||
|
binary.BigEndian.PutUint32(pkt[44:48], 7)
|
||||||
|
binary.BigEndian.PutUint32(pkt[48:52], 99)
|
||||||
|
pkt[52] = 0x50
|
||||||
|
pkt[53] = 0x10 // ACK only
|
||||||
|
binary.BigEndian.PutUint16(pkt[54:56], 65535)
|
||||||
|
|
||||||
|
for i := 0; i < payLen; i++ {
|
||||||
|
pkt[ipLen+tcpLen+i] = byte(i)
|
||||||
|
}
|
||||||
|
return pkt
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestDecodeReadFitsMaxTSOAtDrainThreshold proves the rxBuf sizing is
|
||||||
|
// correct: when rxOff is at the maximum value the drain headroom check
|
||||||
|
// allows, decodeRead must still be able to absorb a worst-case 64KiB
|
||||||
|
// TSO superpacket without dropping the burst. With segmentation deferred
|
||||||
|
// to encrypt time, decodeRead writes only the kernel-supplied bytes into
|
||||||
|
// rxBuf, so the size requirement is just "fit one worst-case input."
|
||||||
|
//
|
||||||
|
// Regression history: in a prior layout the rx buffer doubled as the
|
||||||
|
// segmentation output, a near-threshold drain read returned "scratch too
|
||||||
|
// small", the whole 45-segment TSO burst was dropped, and the remote's TCP
|
||||||
|
// fast-retransmit collapsed cwnd. Keeping this test in the new layout
|
||||||
|
// guards against re-introducing a drain headroom shortfall.
|
||||||
|
func TestDecodeReadFitsMaxTSOAtDrainThreshold(t *testing.T) {
|
||||||
|
const ipv6HdrLen = 40
|
||||||
|
const tcpHdrLen = 20
|
||||||
|
const headerLen = ipv6HdrLen + tcpHdrLen
|
||||||
|
// Maximum TUN read body. The tunReadBufSize cap on readv's body iovec
|
||||||
|
// is what bounds the kernel's superpacket length.
|
||||||
|
pktLen := tunReadBufSize
|
||||||
|
payLen := pktLen - headerLen
|
||||||
|
const targetSegs = 64
|
||||||
|
gsoSize := (payLen + targetSegs - 1) / targetSegs
|
||||||
|
|
||||||
|
pkt := buildTSOv6(payLen, gsoSize)
|
||||||
|
if len(pkt) != pktLen {
|
||||||
|
t.Fatalf("buildTSOv6 produced %d bytes, want %d", len(pkt), pktLen)
|
||||||
|
}
|
||||||
|
|
||||||
|
o := &Offload{
|
||||||
|
rxBuf: make([]byte, tunRxBufCap),
|
||||||
|
}
|
||||||
|
// rxOff at the maximum value the drain headroom check permits before
|
||||||
|
// it would refuse another read. Any drain-time read up to this
|
||||||
|
// threshold MUST still process correctly.
|
||||||
|
o.rxOff = tunRxBufCap - tunRxBufSize
|
||||||
|
|
||||||
|
// Stage the body in rxBuf as if readv(2) just placed it there.
|
||||||
|
copy(o.rxBuf[o.rxOff:], pkt)
|
||||||
|
|
||||||
|
// Encode the matching virtio_net_hdr.
|
||||||
|
hdr := virtio.Hdr{
|
||||||
|
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||||
|
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV6,
|
||||||
|
HdrLen: uint16(headerLen),
|
||||||
|
GSOSize: uint16(gsoSize),
|
||||||
|
CsumStart: uint16(ipv6HdrLen),
|
||||||
|
CsumOffset: 16,
|
||||||
|
}
|
||||||
|
hdr.Encode(o.readVnetScratch[:])
|
||||||
|
|
||||||
|
startRxOff := o.rxOff
|
||||||
|
if err := o.decodeRead(pktLen); err != nil {
|
||||||
|
t.Fatalf("decodeRead at drain threshold returned %v — rxBuf sizing regression: "+
|
||||||
|
"tunRxBufSize=%d must hold one worst-case input (%d)",
|
||||||
|
err, tunRxBufSize, pktLen)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(o.pending) != 1 {
|
||||||
|
t.Fatalf("got %d packets, want 1 superpacket entry", len(o.pending))
|
||||||
|
}
|
||||||
|
got := o.pending[0]
|
||||||
|
if !got.GSO.IsSuperpacket() {
|
||||||
|
t.Fatalf("expected superpacket GSO metadata, got %+v", got.GSO)
|
||||||
|
}
|
||||||
|
if got.GSO.Proto != GSOProtoTCP {
|
||||||
|
t.Errorf("GSO.Proto=%d want TCP", got.GSO.Proto)
|
||||||
|
}
|
||||||
|
if got.GSO.Size != uint16(gsoSize) {
|
||||||
|
t.Errorf("GSO.Size=%d want %d", got.GSO.Size, gsoSize)
|
||||||
|
}
|
||||||
|
if got.GSO.HdrLen != uint16(headerLen) {
|
||||||
|
t.Errorf("GSO.HdrLen=%d want %d", got.GSO.HdrLen, headerLen)
|
||||||
|
}
|
||||||
|
if got.GSO.CsumStart != uint16(ipv6HdrLen) {
|
||||||
|
t.Errorf("GSO.CsumStart=%d want %d", got.GSO.CsumStart, ipv6HdrLen)
|
||||||
|
}
|
||||||
|
if len(got.Bytes) != pktLen {
|
||||||
|
t.Errorf("len(Bytes)=%d want %d", len(got.Bytes), pktLen)
|
||||||
|
}
|
||||||
|
|
||||||
|
// rxOff advances exactly by the kernel-supplied body length — no
|
||||||
|
// segmentation output to account for any more.
|
||||||
|
if o.rxOff != startRxOff+pktLen {
|
||||||
|
t.Errorf("rxOff=%d want %d", o.rxOff, startRxOff+pktLen)
|
||||||
|
}
|
||||||
|
if o.rxOff > tunRxBufCap {
|
||||||
|
t.Fatalf("rxOff=%d overran rxBuf (cap=%d)", o.rxOff, tunRxBufCap)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate that segmenting the returned superpacket reproduces the
|
||||||
|
// expected per-segment IPv6 payload length and TCP checksum.
|
||||||
|
wantSegs := (payLen + gsoSize - 1) / gsoSize
|
||||||
|
gotSegs := 0
|
||||||
|
if err := SegmentSuperpacket(got, func(seg []byte) error {
|
||||||
|
defer func() { gotSegs++ }()
|
||||||
|
if len(seg) < headerLen+1 {
|
||||||
|
t.Errorf("seg %d too short: %d", gotSegs, len(seg))
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if seg[0]>>4 != 6 {
|
||||||
|
t.Errorf("seg %d: bad IP version %#x", gotSegs, seg[0])
|
||||||
|
}
|
||||||
|
segPay := len(seg) - headerLen
|
||||||
|
gotPL := binary.BigEndian.Uint16(seg[4:6])
|
||||||
|
if gotPL != uint16(tcpHdrLen+segPay) {
|
||||||
|
t.Errorf("seg %d: payload_len=%d want %d", gotSegs, gotPL, tcpHdrLen+segPay)
|
||||||
|
}
|
||||||
|
psum := pseudoHeaderIPv6(seg[8:24], seg[24:40], unix.IPPROTO_TCP, tcpHdrLen+segPay)
|
||||||
|
if !verifyChecksum(seg[ipv6HdrLen:], psum) {
|
||||||
|
t.Errorf("seg %d: bad TCP checksum", gotSegs)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("SegmentSuperpacket: %v", err)
|
||||||
|
}
|
||||||
|
if gotSegs != wantSegs {
|
||||||
|
t.Fatalf("got %d segments, want %d", gotSegs, wantSegs)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,43 @@
|
|||||||
|
//go:build linux && !android
|
||||||
|
// +build linux,!android
|
||||||
|
|
||||||
|
package virtio
|
||||||
|
|
||||||
|
import "encoding/binary"
|
||||||
|
|
||||||
|
// Size is the on-wire length of struct virtio_net_hdr the kernel
|
||||||
|
// prepends/expects on a TUN opened with IFF_VNET_HDR (TUNSETVNETHDRSZ
|
||||||
|
// not set).
|
||||||
|
const Size = 10
|
||||||
|
|
||||||
|
// Hdr is the Go view of the legacy virtio_net_hdr.
|
||||||
|
type Hdr struct {
|
||||||
|
Flags uint8
|
||||||
|
GSOType uint8
|
||||||
|
HdrLen uint16
|
||||||
|
GSOSize uint16
|
||||||
|
CsumStart uint16
|
||||||
|
CsumOffset uint16
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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])
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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) {
|
||||||
|
b[0] = h.Flags
|
||||||
|
b[1] = h.GSOType
|
||||||
|
binary.NativeEndian.PutUint16(b[2:4], h.HdrLen)
|
||||||
|
binary.NativeEndian.PutUint16(b[4:6], h.GSOSize)
|
||||||
|
binary.NativeEndian.PutUint16(b[6:8], h.CsumStart)
|
||||||
|
binary.NativeEndian.PutUint16(b[8:10], h.CsumOffset)
|
||||||
|
}
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
//go:build !linux || android
|
||||||
|
|
||||||
|
package virtio
|
||||||
@@ -0,0 +1,434 @@
|
|||||||
|
//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; the array is sized to that
|
||||||
|
// worst case so the snapshot lives on the stack with no per-call heap
|
||||||
|
// allocation.
|
||||||
|
const maxSegHdrLen = ipv4HeaderMaxLen + tcpHeaderMaxLen // 120
|
||||||
|
|
||||||
|
// Byte offsets inside an IPv4 header.
|
||||||
|
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
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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 {
|
||||||
|
// When RSC_INFO is set the csum_start/csum_offset fields are repurposed to
|
||||||
|
// carry coalescing info rather than checksum offsets. A TUN writing via
|
||||||
|
// IFF_VNET_HDR should never emit this, but if it did we would silently
|
||||||
|
// miscompute the segment checksums — refuse the packet instead.
|
||||||
|
if hdr.Flags&unix.VIRTIO_NET_HDR_F_RSC_INFO != 0 {
|
||||||
|
return fmt.Errorf("virtio RSC_INFO flag not supported on TUN reads")
|
||||||
|
}
|
||||||
|
if len(pkt) < ipv4HeaderMinLen {
|
||||||
|
return fmt.Errorf("packet too short")
|
||||||
|
}
|
||||||
|
ipVersion := pkt[0] >> 4
|
||||||
|
switch hdr.GSOType {
|
||||||
|
case unix.VIRTIO_NET_HDR_GSO_TCPV4:
|
||||||
|
if ipVersion != 4 {
|
||||||
|
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
||||||
|
}
|
||||||
|
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 < 20 || tcpHLen > 60 {
|
||||||
|
// A TCP header must be between 20 and 60 bytes in length.
|
||||||
|
return fmt.Errorf("tcp header len is invalid: %d", tcpHLen)
|
||||||
|
}
|
||||||
|
hdr.HdrLen = hdr.CsumStart + tcpHLen
|
||||||
|
}
|
||||||
|
|
||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
// SegmentTCP walks a TSO superpacket pkt, yielding each segment as a
|
||||||
|
// slice into pkt itself. Per-segment plaintext is laid out by stamping a
|
||||||
|
// copy of the original L3+L4 header into pkt at offset i*gsoSize, where it
|
||||||
|
// sits immediately before that segment's payload chunk in the original
|
||||||
|
// buffer. The stamp is destructive but harmless: iter i's header write lands
|
||||||
|
// on pkt[i*G : i*G+hdrLen], which is the tail of seg_{i-1}'s payload (already
|
||||||
|
// consumed) and ends exactly where seg_i's payload begins, so it never clobbers
|
||||||
|
// live payload — this holds even when gsoSize < hdrLen. The header bytes are
|
||||||
|
// sourced from a pristine snapshot taken before the loop (savedHdr), NOT from
|
||||||
|
// pkt[:hdrLen], because when gsoSize < hdrLen the stamps would otherwise
|
||||||
|
// overwrite the leading header in place and every stamp after the first would
|
||||||
|
// copy corrupted bytes. pkt is consumed by this call and must not be inspected
|
||||||
|
// by the caller after the final yield.
|
||||||
|
func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
|
||||||
|
if gsoSizeU == 0 {
|
||||||
|
return fmt.Errorf("gso_size is zero")
|
||||||
|
}
|
||||||
|
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 := (payLen + gsoSize - 1) / gsoSize
|
||||||
|
if numSeg == 0 {
|
||||||
|
numSeg = 1
|
||||||
|
}
|
||||||
|
|
||||||
|
origSeq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
|
||||||
|
origFlags := pkt[csumStart+tcpFlagsOff]
|
||||||
|
|
||||||
|
var tmp [tcpHeaderMaxLen]byte
|
||||||
|
copy(tmp[:tcpHdrLen], pkt[csumStart:headerLen])
|
||||||
|
tmp[tcpSeqOff], tmp[tcpSeqOff+1], tmp[tcpSeqOff+2], tmp[tcpSeqOff+3] = 0, 0, 0, 0
|
||||||
|
tmp[tcpFlagsOff] = 0
|
||||||
|
tmp[tcpChecksumOff], tmp[tcpChecksumOff+1] = 0, 0
|
||||||
|
baseTcpHdrSum := uint32(checksum.Checksum(tmp[:tcpHdrLen], 0))
|
||||||
|
|
||||||
|
var baseProtoSum uint32
|
||||||
|
if isV4 {
|
||||||
|
baseProtoSum = uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0))
|
||||||
|
} else {
|
||||||
|
baseProtoSum = uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0))
|
||||||
|
}
|
||||||
|
baseProtoSum += uint32(unix.IPPROTO_TCP)
|
||||||
|
|
||||||
|
var origIPID uint16
|
||||||
|
var baseIPHdrSum uint32
|
||||||
|
if isV4 {
|
||||||
|
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
|
||||||
|
ihl := int(pkt[0]&0x0f) * 4
|
||||||
|
if ihl < ipv4HeaderMinLen || ihl > csumStart {
|
||||||
|
return fmt.Errorf("bad IPv4 IHL: %d", ihl)
|
||||||
|
}
|
||||||
|
var ipTmp [ipv4HeaderMaxLen]byte
|
||||||
|
copy(ipTmp[:ihl], pkt[:ihl])
|
||||||
|
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
|
||||||
|
ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0
|
||||||
|
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
|
||||||
|
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Snapshot the pristine L3+L4 header once. Every segment's header is
|
||||||
|
// stamped from this copy, so overlapping stamps (gsoSize < headerLen)
|
||||||
|
// can never corrupt the source. The variable fields (seq/flags/cksum/
|
||||||
|
// totalLen/id) captured here are stale but are overwritten per segment.
|
||||||
|
var savedHdr [maxSegHdrLen]byte
|
||||||
|
copy(savedHdr[:headerLen], pkt[:headerLen])
|
||||||
|
|
||||||
|
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 pristine snapshot. Iter 0's header is
|
||||||
|
// already at pkt[:headerLen] (identical to savedHdr), so only i ≥ 1
|
||||||
|
// needs the stamp. The per-segment patches below overwrite the
|
||||||
|
// variable fields.
|
||||||
|
if i > 0 {
|
||||||
|
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*G : i*G+H], which is
|
||||||
|
// the tail of seg_{i-1}'s payload (already consumed) and never
|
||||||
|
// overlaps seg_i's own payload at pkt[H+i*G : H+(i+1)*G].
|
||||||
|
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
|
||||||
|
wide := uint64(baseTcpHdrSum) + uint64(paySum) + uint64(baseProtoSum)
|
||||||
|
wide += uint64(segSeq) + uint64(segFlags) + uint64(tcpLen)
|
||||||
|
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*G:i*G+segLen] to the caller. Per-segment patches are total_len +
|
||||||
|
// IPv4 csum (or IPv6 payload_len) plus the UDP length and checksum. pkt is
|
||||||
|
// consumed destructively; see SegmentTCP for the layout reasoning, including
|
||||||
|
// why the header is stamped from a pristine snapshot rather than pkt[:hdrLen]
|
||||||
|
// (correctness when gsoSize < hdrLen).
|
||||||
|
//
|
||||||
|
// UDP-GSO leaves the IPv4 ID identical across segments (the kernel does not
|
||||||
|
// bump it), which is why the IP-level per-segment work is limited to
|
||||||
|
// total_len + IPv4 header checksum (v4) or payload_len (v6).
|
||||||
|
func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
|
||||||
|
if gsoSizeU == 0 {
|
||||||
|
return fmt.Errorf("gso_size is zero")
|
||||||
|
}
|
||||||
|
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 := (payLen + gsoSize - 1) / gsoSize
|
||||||
|
if numSeg == 0 {
|
||||||
|
numSeg = 1
|
||||||
|
}
|
||||||
|
|
||||||
|
var udpTmp [udpHeaderLen]byte
|
||||||
|
copy(udpTmp[:], pkt[csumStart:headerLen])
|
||||||
|
udpTmp[udpLengthOff], udpTmp[udpLengthOff+1] = 0, 0
|
||||||
|
udpTmp[udpChecksumOff], udpTmp[udpChecksumOff+1] = 0, 0
|
||||||
|
baseUDPHdrSum := uint32(checksum.Checksum(udpTmp[:], 0))
|
||||||
|
|
||||||
|
var baseProtoSum uint32
|
||||||
|
if isV4 {
|
||||||
|
baseProtoSum = uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0))
|
||||||
|
} else {
|
||||||
|
baseProtoSum = uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0))
|
||||||
|
}
|
||||||
|
baseProtoSum += uint32(unix.IPPROTO_UDP)
|
||||||
|
|
||||||
|
var baseIPHdrSum uint32
|
||||||
|
if isV4 {
|
||||||
|
ihl := int(pkt[0]&0x0f) * 4
|
||||||
|
if ihl < ipv4HeaderMinLen || ihl > csumStart {
|
||||||
|
return fmt.Errorf("bad IPv4 IHL: %d", ihl)
|
||||||
|
}
|
||||||
|
var ipTmp [ipv4HeaderMaxLen]byte
|
||||||
|
copy(ipTmp[:ihl], pkt[:ihl])
|
||||||
|
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
|
||||||
|
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
|
||||||
|
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Snapshot the pristine L3+L4 header once and stamp every segment from
|
||||||
|
// it; see SegmentTCP for why sourcing from pkt[:headerLen] corrupts
|
||||||
|
// segments when gsoSize < headerLen.
|
||||||
|
var savedHdr [maxSegHdrLen]byte
|
||||||
|
copy(savedHdr[:headerLen], pkt[:headerLen])
|
||||||
|
|
||||||
|
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 {
|
||||||
|
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
|
||||||
|
ipSum := baseIPHdrSum + uint32(totalLen)
|
||||||
|
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
|
||||||
|
} else {
|
||||||
|
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
|
||||||
|
}
|
||||||
|
|
||||||
|
binary.BigEndian.PutUint16(seg[csumStart+udpLengthOff:csumStart+udpLengthOff+2], uint16(udpLen))
|
||||||
|
|
||||||
|
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
|
||||||
|
wide := uint64(baseUDPHdrSum) + uint64(paySum) + uint64(baseProtoSum)
|
||||||
|
wide += uint64(udpLen) + uint64(udpLen)
|
||||||
|
wide = (wide & 0xffffffff) + (wide >> 32)
|
||||||
|
wide = (wide & 0xffffffff) + (wide >> 32)
|
||||||
|
csum := foldComplement(uint32(wide))
|
||||||
|
if csum == 0 {
|
||||||
|
csum = 0xffff
|
||||||
|
}
|
||||||
|
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. csum_start / csum_offset point at the 16-bit
|
||||||
|
// checksum field; we zero it, fold a full sum (the field was pre-loaded with
|
||||||
|
// the pseudo-header partial sum by the kernel), and store the result.
|
||||||
|
func FinishChecksum(seg []byte, hdr Hdr) error {
|
||||||
|
cs := int(hdr.CsumStart)
|
||||||
|
co := int(hdr.CsumOffset)
|
||||||
|
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
|
||||||
|
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], ^checksum.Checksum(seg[cs:], partial))
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// foldComplement folds a 32-bit one's-complement partial sum to 16 bits and
|
||||||
|
// complements it, yielding the on-wire Internet checksum value.
|
||||||
|
func foldComplement(sum uint32) uint16 {
|
||||||
|
sum = (sum & 0xffff) + (sum >> 16)
|
||||||
|
sum = (sum & 0xffff) + (sum >> 16)
|
||||||
|
return ^uint16(sum)
|
||||||
|
}
|
||||||
@@ -0,0 +1,335 @@
|
|||||||
|
//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 := Hdr{
|
||||||
|
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||||
|
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
|
||||||
|
GSOSize: 6, // two 6-byte segments
|
||||||
|
CsumStart: csumStart,
|
||||||
|
CsumOffset: 6,
|
||||||
|
}
|
||||||
|
if err := CorrectHdrLen(pkt, &hdr); err != nil {
|
||||||
|
t.Fatalf("CorrectHdrLen rejected a valid 40-byte USO superpacket: %v", err)
|
||||||
|
}
|
||||||
|
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 := Hdr{
|
||||||
|
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||||
|
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
|
||||||
|
GSOSize: 6,
|
||||||
|
CsumStart: 20,
|
||||||
|
CsumOffset: 6,
|
||||||
|
}
|
||||||
|
if err := CorrectHdrLen(pkt, &hdr); err == nil {
|
||||||
|
t.Fatalf("CorrectHdrLen accepted a too-short (25-byte) packet")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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)
|
||||||
|
}
|
||||||
|
// UDP-GSO keeps the same IPv4 ID across every segment.
|
||||||
|
if id := binary.BigEndian.Uint16(seg[4:6]); id != 0x4242 {
|
||||||
|
t.Errorf("seg %d: ip id=%#x want 0x4242", i, id)
|
||||||
|
}
|
||||||
|
|
||||||
|
segPayLen := len(seg) - int(hdrLen)
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -550,7 +550,7 @@ func (t *tun) Read(to []byte) (int, error) {
|
|||||||
return n - 4, nil
|
return n - 4, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Write pushes one IP packet onto the utun device.
|
// Write pushes one IP packet onto the utun device. Only valid for single threaded use.
|
||||||
func (t *tun) Write(from []byte) (int, error) {
|
func (t *tun) Write(from []byte) (int, error) {
|
||||||
if len(from) == 0 {
|
if len(from) == 0 {
|
||||||
return 0, syscall.EIO
|
return 0, syscall.EIO
|
||||||
|
|||||||
+129
-8
@@ -34,6 +34,23 @@ type tun struct {
|
|||||||
TXQueueLen int
|
TXQueueLen int
|
||||||
deviceIndex int
|
deviceIndex int
|
||||||
ioctlFd uintptr
|
ioctlFd uintptr
|
||||||
|
vnetHdr bool
|
||||||
|
// offloadFlags is the exact TUN_F_* offload mask newTun negotiated with
|
||||||
|
// the kernel: usoOffloadFlags when USO was accepted, tsoOffloadFlags on
|
||||||
|
// the TSO-only fallback, or 0 when vnetHdr is off. TUNSETOFFLOAD is
|
||||||
|
// device-wide (drivers/net/tun.c set_offload updates tun->set_features
|
||||||
|
// for the whole netdev), so addQueue must replay this exact
|
||||||
|
// mask on every added queue — issuing a narrower mask there would
|
||||||
|
// silently downgrade offloads (e.g. disable USO) for all queues while
|
||||||
|
// they still advertise the stale capability.
|
||||||
|
offloadFlags uint
|
||||||
|
// routeFeatureECN, when true, sets RTAX_FEATURE_ECN on every route we
|
||||||
|
// install for the tun. The kernel then actively negotiates ECN for
|
||||||
|
// connections destined to those prefixes (equivalent to `ip route
|
||||||
|
// change ... features ecn`) regardless of net.ipv4.tcp_ecn, so flows
|
||||||
|
// across the nebula mesh use ECN even when the host default is the
|
||||||
|
// passive setting (=2). Disable via tunnels.ecn=false.
|
||||||
|
routeFeatureECN bool
|
||||||
|
|
||||||
Routes atomic.Pointer[[]Route]
|
Routes atomic.Pointer[[]Route]
|
||||||
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
||||||
@@ -72,7 +89,9 @@ type ifreqQLEN struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
t, err := newTunGeneric(c, l, deviceFd, vpnNetworks)
|
// We don't know what flags the caller opened this fd with and can't turn
|
||||||
|
// on IFF_VNET_HDR after TUNSETIFF, so skip offload on inherited fds.
|
||||||
|
t, err := newTunGeneric(c, l, deviceFd, false, 0, vpnNetworks)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -117,6 +136,26 @@ func tunSetIff(fd int, name string, flags uint16) (string, error) {
|
|||||||
return strings.Trim(string(req.Name[:]), "\x00"), nil
|
return strings.Trim(string(req.Name[:]), "\x00"), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// tsoOffloadFlags are the TUN_F_* bits we ask the kernel to enable when a
|
||||||
|
// TSO-capable TUN is available. CSUM is required as a prerequisite for TSO.
|
||||||
|
// TSO_ECN tells the kernel we propagate ECN correctly through coalesce and
|
||||||
|
// segmentation, so it can deliver superpackets whose seed has CWR/ECE set
|
||||||
|
// or whose IP-level codepoint is CE.
|
||||||
|
const tsoOffloadFlags = unix.TUN_F_CSUM | unix.TUN_F_TSO4 | unix.TUN_F_TSO6 | unix.TUN_F_TSO_ECN
|
||||||
|
|
||||||
|
// usoOffloadFlags adds UDP Segmentation Offload to tsoOffloadFlags. Requires
|
||||||
|
// Linux ≥ 6.2; older kernels reject it and we fall back to TCP-only TSO via
|
||||||
|
// tsoOffloadFlags.
|
||||||
|
const usoOffloadFlags = tsoOffloadFlags | unix.TUN_F_USO4 | unix.TUN_F_USO6
|
||||||
|
|
||||||
|
// offloadUSOEnabled reports whether the negotiated offload mask includes UDP
|
||||||
|
// Segmentation Offload. It is the single source of truth for the usoEnabled
|
||||||
|
// capability surfaced by each queue, so the mask stored on the tun and the USO
|
||||||
|
// bit reported to coalescers can never drift apart.
|
||||||
|
func offloadUSOEnabled(offloadFlags uint) bool {
|
||||||
|
return offloadFlags&(unix.TUN_F_USO4|unix.TUN_F_USO6) != 0
|
||||||
|
}
|
||||||
|
|
||||||
func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) {
|
func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) {
|
||||||
baseFlags := uint16(unix.IFF_TUN | unix.IFF_NO_PI)
|
baseFlags := uint16(unix.IFF_TUN | unix.IFF_NO_PI)
|
||||||
if multiqueue {
|
if multiqueue {
|
||||||
@@ -124,17 +163,56 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
|
|||||||
}
|
}
|
||||||
nameStr := c.GetString("tun.dev", "")
|
nameStr := c.GetString("tun.dev", "")
|
||||||
|
|
||||||
|
// First try to enable IFF_VNET_HDR via TUNSETIFF and negotiate TUN_F_*
|
||||||
|
// offloads via TUNSETOFFLOAD so we can receive TSO/USO superpackets.
|
||||||
|
// We try TSO+USO first, fall back to TSO-only on kernels without USO
|
||||||
|
// (Linux < 6.2), and finally give up on virtio headers entirely and
|
||||||
|
// reopen as a plain TUN if neither offload mask is accepted.
|
||||||
fd, err := openTunDev()
|
fd, err := openTunDev()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
name, err := tunSetIff(fd, nameStr, baseFlags)
|
vnetHdr := true
|
||||||
|
// offloadFlags is the exact TUN_F_* mask the kernel accepted. We remember
|
||||||
|
// it (rather than a plain bool) so addQueue can replay the
|
||||||
|
// identical device-wide mask on added queues instead of downgrading them.
|
||||||
|
var offloadFlags uint
|
||||||
|
name, err := tunSetIff(fd, nameStr, baseFlags|unix.IFF_VNET_HDR)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = unix.Close(fd)
|
_ = unix.Close(fd)
|
||||||
return nil, &NameError{Name: nameStr, Underlying: err}
|
vnetHdr = false
|
||||||
|
} else {
|
||||||
|
// Try TSO+USO first. On kernels without USO support (Linux < 6.2)
|
||||||
|
// the ioctl returns EINVAL; fall back to the TCP-only mask before
|
||||||
|
// giving up on VNET_HDR entirely.
|
||||||
|
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(usoOffloadFlags)); err == nil {
|
||||||
|
offloadFlags = usoOffloadFlags
|
||||||
|
} else if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(tsoOffloadFlags)); err == nil {
|
||||||
|
offloadFlags = tsoOffloadFlags
|
||||||
|
} else {
|
||||||
|
l.Warn("Failed to enable TUN offload (TSO); proceeding without virtio headers", "error", err)
|
||||||
|
_ = unix.Close(fd)
|
||||||
|
vnetHdr = false
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
t, err := newTunGeneric(c, l, fd, vpnNetworks)
|
if !vnetHdr {
|
||||||
|
fd, err = openTunDev()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
name, err = tunSetIff(fd, nameStr, baseFlags)
|
||||||
|
if err != nil {
|
||||||
|
_ = unix.Close(fd)
|
||||||
|
return nil, &NameError{Name: nameStr, Underlying: err}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if vnetHdr {
|
||||||
|
l.Info("TUN offload enabled", "tso", true, "uso", offloadUSOEnabled(offloadFlags))
|
||||||
|
}
|
||||||
|
|
||||||
|
t, err := newTunGeneric(c, l, fd, vnetHdr, offloadFlags, vpnNetworks)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -145,9 +223,19 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
|
|||||||
}
|
}
|
||||||
|
|
||||||
// newTunGeneric does all the stuff common to different tun initialization
|
// newTunGeneric does all the stuff common to different tun initialization
|
||||||
// paths. It will close your files on error.
|
// paths. It will close your files on error. offloadFlags is the TUN_F_* mask
|
||||||
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
// newTun negotiated (0 when vnetHdr is off); the queues' USO capability is
|
||||||
qs, err := tio.NewPollQueueSet()
|
// derived from it so it can never disagree with the mask we replay on added
|
||||||
|
// multiqueue readers.
|
||||||
|
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFlags uint, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
|
var qs tio.QueueSet
|
||||||
|
var err error
|
||||||
|
if vnetHdr {
|
||||||
|
qs, err = tio.NewOffloadQueueSet(offloadUSOEnabled(offloadFlags))
|
||||||
|
} else {
|
||||||
|
qs, err = tio.NewPollQueueSet()
|
||||||
|
}
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = unix.Close(fd)
|
_ = unix.Close(fd)
|
||||||
return nil, err
|
return nil, err
|
||||||
@@ -161,10 +249,13 @@ func newTunGeneric(c *config.C, l *slog.Logger, fd int, vpnNetworks []netip.Pref
|
|||||||
t := &tun{
|
t := &tun{
|
||||||
readers: qs,
|
readers: qs,
|
||||||
closeLock: sync.Mutex{},
|
closeLock: sync.Mutex{},
|
||||||
|
vnetHdr: vnetHdr,
|
||||||
|
offloadFlags: offloadFlags,
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
TXQueueLen: c.GetInt("tun.tx_queue", 500),
|
TXQueueLen: c.GetInt("tun.tx_queue", 500),
|
||||||
useSystemRoutes: c.GetBool("tun.use_system_route_table", false),
|
useSystemRoutes: c.GetBool("tun.use_system_route_table", false),
|
||||||
useSystemRoutesBufferSize: c.GetInt("tun.use_system_route_table_buffer_size", 0),
|
useSystemRoutesBufferSize: c.GetInt("tun.use_system_route_table_buffer_size", 0),
|
||||||
|
routeFeatureECN: c.GetBool("tunnels.ecn", true),
|
||||||
routesFromSystem: map[netip.Prefix]routing.Gateways{},
|
routesFromSystem: map[netip.Prefix]routing.Gateways{},
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
@@ -259,7 +350,8 @@ func (t *tun) reload(c *config.C, initial bool) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Queues opens additional kernel multiqueue fds until the device has n
|
// Queues opens additional kernel multiqueue fds until the device has n
|
||||||
// queues, then returns them all. The first queue was opened by newTun.
|
// queues, then returns them all. The first queue was opened by newTun; each
|
||||||
|
// extra fd replays the negotiated offload state (see addQueue).
|
||||||
func (t *tun) Queues(n int) ([]tio.Queue, error) {
|
func (t *tun) Queues(n int) ([]tio.Queue, error) {
|
||||||
for len(t.readers.Queues()) < n {
|
for len(t.readers.Queues()) < n {
|
||||||
if err := t.addQueue(); err != nil {
|
if err := t.addQueue(); err != nil {
|
||||||
@@ -281,11 +373,25 @@ func (t *tun) addQueue() error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
flags := uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_MULTI_QUEUE)
|
flags := uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_MULTI_QUEUE)
|
||||||
|
if t.vnetHdr {
|
||||||
|
flags |= unix.IFF_VNET_HDR
|
||||||
|
}
|
||||||
if _, err = tunSetIff(fd, t.Device, flags); err != nil {
|
if _, err = tunSetIff(fd, t.Device, flags); err != nil {
|
||||||
_ = unix.Close(fd)
|
_ = unix.Close(fd)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if t.vnetHdr {
|
||||||
|
// Replay the exact mask newTun negotiated. TUNSETOFFLOAD is
|
||||||
|
// device-wide, so issuing the TSO-only mask here would disable USO
|
||||||
|
// for every queue (including queue 0) on kernels where newTun
|
||||||
|
// successfully enabled it, while the queues keep advertising USO.
|
||||||
|
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(t.offloadFlags)); err != nil {
|
||||||
|
_ = unix.Close(fd)
|
||||||
|
return fmt.Errorf("failed to enable offload on multiqueue tun fd: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
err = t.readers.Add(fd)
|
err = t.readers.Add(fd)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = unix.Close(fd)
|
_ = unix.Close(fd)
|
||||||
@@ -460,6 +566,18 @@ func (t *tun) setDefaultRoute(cidr netip.Prefix) error {
|
|||||||
Table: unix.RT_TABLE_MAIN,
|
Table: unix.RT_TABLE_MAIN,
|
||||||
Type: unix.RTN_UNICAST,
|
Type: unix.RTN_UNICAST,
|
||||||
}
|
}
|
||||||
|
// Match the metric the kernel uses for its auto-installed connected
|
||||||
|
// route, so RouteReplace overwrites it in place instead of adding a
|
||||||
|
// second route at a worse metric. IPv6 connected routes are installed
|
||||||
|
// at metric 256 (IP6_RT_PRIO_KERN); IPv4 uses 0. Without this, the
|
||||||
|
// kernel route wins lookups and our MTU / AdvMSS / Features never
|
||||||
|
// apply on v6.
|
||||||
|
if cidr.Addr().Is6() {
|
||||||
|
nr.Priority = 256
|
||||||
|
}
|
||||||
|
if t.routeFeatureECN {
|
||||||
|
nr.Features |= unix.RTAX_FEATURE_ECN
|
||||||
|
}
|
||||||
err := netlink.RouteReplace(&nr)
|
err := netlink.RouteReplace(&nr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr)
|
t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr)
|
||||||
@@ -509,6 +627,9 @@ func (t *tun) addRoutes(logErrors bool) error {
|
|||||||
if r.Metric > 0 {
|
if r.Metric > 0 {
|
||||||
nr.Priority = r.Metric
|
nr.Priority = r.Metric
|
||||||
}
|
}
|
||||||
|
if t.routeFeatureECN {
|
||||||
|
nr.Features |= unix.RTAX_FEATURE_ECN
|
||||||
|
}
|
||||||
|
|
||||||
err := netlink.RouteReplace(&nr)
|
err := netlink.RouteReplace(&nr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -3,7 +3,9 @@
|
|||||||
|
|
||||||
package overlay
|
package overlay
|
||||||
|
|
||||||
import "testing"
|
import (
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
var runAdvMSSTests = []struct {
|
var runAdvMSSTests = []struct {
|
||||||
name string
|
name string
|
||||||
@@ -32,3 +34,65 @@ func TestTunAdvMSS(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestOffloadUSOEnabled pins the single source of truth for the per-queue USO
|
||||||
|
// capability: it is derived from the negotiated offload mask, so the mask
|
||||||
|
// stored on the tun and the capability reported to coalescers cannot drift.
|
||||||
|
func TestOffloadUSOEnabled(t *testing.T) {
|
||||||
|
// usoOffloadFlags must be a strict superset of tsoOffloadFlags. Otherwise
|
||||||
|
// the TSO-only fallback (and the historic hardcoded-mask bug in
|
||||||
|
// addQueue) would not actually be a downgrade.
|
||||||
|
if usoOffloadFlags&tsoOffloadFlags != tsoOffloadFlags {
|
||||||
|
t.Fatalf("usoOffloadFlags (%#x) is not a superset of tsoOffloadFlags (%#x)", usoOffloadFlags, tsoOffloadFlags)
|
||||||
|
}
|
||||||
|
if usoOffloadFlags == tsoOffloadFlags {
|
||||||
|
t.Fatal("usoOffloadFlags must add bits beyond tsoOffloadFlags")
|
||||||
|
}
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
offloadFlags uint
|
||||||
|
wantUSO bool
|
||||||
|
}{
|
||||||
|
{"uso-negotiated", usoOffloadFlags, true},
|
||||||
|
{"tso-fallback", tsoOffloadFlags, false},
|
||||||
|
{"no-vnet-hdr", 0, false},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
if got := offloadUSOEnabled(tc.offloadFlags); got != tc.wantUSO {
|
||||||
|
t.Fatalf("offloadUSOEnabled(%#x) = %v, want %v", tc.offloadFlags, got, tc.wantUSO)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestAddQueueReplaysNegotiatedMask guards the device-wide TUNSETOFFLOAD
|
||||||
|
// downgrade bug: addQueue must issue the exact mask newTun negotiated
|
||||||
|
// (t.offloadFlags), not a hardcoded TSO-only mask. Because TUNSETOFFLOAD is
|
||||||
|
// per-netdev, a narrower mask on an added queue silently disables USO for
|
||||||
|
// every queue on a USO-capable kernel while the queues keep advertising it.
|
||||||
|
//
|
||||||
|
// A full multi-queue exercise needs /dev/net/tun and CAP_NET_ADMIN, which are
|
||||||
|
// not available in CI/sandbox, so this asserts on the struct field that the
|
||||||
|
// TUNSETOFFLOAD argument is read from.
|
||||||
|
func TestAddQueueReplaysNegotiatedMask(t *testing.T) {
|
||||||
|
t.Run("uso-negotiated", func(t *testing.T) {
|
||||||
|
tn := &tun{vnetHdr: true, offloadFlags: usoOffloadFlags}
|
||||||
|
// The ioctl argument in addQueue is uintptr(t.offloadFlags);
|
||||||
|
// it must equal the negotiated USO mask, and must NOT be the TSO-only
|
||||||
|
// mask (the original bug).
|
||||||
|
if tn.offloadFlags != usoOffloadFlags {
|
||||||
|
t.Fatalf("offloadFlags = %#x, want %#x", tn.offloadFlags, usoOffloadFlags)
|
||||||
|
}
|
||||||
|
if tn.offloadFlags == tsoOffloadFlags {
|
||||||
|
t.Fatal("added queue would downgrade USO: offloadFlags must not be the TSO-only mask when USO was negotiated")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
t.Run("tso-fallback", func(t *testing.T) {
|
||||||
|
tn := &tun{vnetHdr: true, offloadFlags: tsoOffloadFlags}
|
||||||
|
if tn.offloadFlags != tsoOffloadFlags {
|
||||||
|
t.Fatalf("offloadFlags = %#x, want %#x", tn.offloadFlags, tsoOffloadFlags)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,33 @@
|
|||||||
|
//go:build debug
|
||||||
|
|
||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"log/slog"
|
||||||
|
"net/http"
|
||||||
|
_ "net/http/pprof" // registers pprof handlers on http.DefaultServeMux
|
||||||
|
)
|
||||||
|
|
||||||
|
// startPprofServer serves net/http/pprof on :6060 for the life of ctx. It is
|
||||||
|
// only compiled into debug builds (`-tags debug`, `make debug`), so a debug
|
||||||
|
// build announces itself with the Info line below.
|
||||||
|
func startPprofServer(ctx context.Context, l *slog.Logger) {
|
||||||
|
server := &http.Server{Addr: ":6060", Handler: nil}
|
||||||
|
l.Info("Starting pprof debug server (debug build)", "addr", server.Addr)
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
if err := server.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
|
||||||
|
l.Error("pprof debug server stopped", "error", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Shut down the server when the context is cancelled.
|
||||||
|
go func() {
|
||||||
|
<-ctx.Done()
|
||||||
|
if err := server.Shutdown(context.Background()); err != nil {
|
||||||
|
l.Debug("Error shutting down pprof debug server", "error", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
//go:build !debug
|
||||||
|
|
||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"log/slog"
|
||||||
|
)
|
||||||
|
|
||||||
|
// startPprofServer is a no-op unless built with `-tags debug` (see make debug).
|
||||||
|
func startPprofServer(_ context.Context, _ *slog.Logger) {}
|
||||||
+20
-2
@@ -4,6 +4,8 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"testing"
|
"testing"
|
||||||
@@ -89,7 +91,23 @@ func newSimpleService(caCrt cert.Certificate, caKey []byte, name string, udpIp n
|
|||||||
return s
|
return s
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ephemeralUDPPort reserves a free UDP port by binding port 0, then releases
|
||||||
|
// it for the caller to use. A fixed port would collide with concurrent test
|
||||||
|
// runs or an unrelated process (say, a real nebula) already listening on it.
|
||||||
|
func ephemeralUDPPort(t *testing.T) int {
|
||||||
|
t.Helper()
|
||||||
|
pc, err := net.ListenPacket("udp4", "127.0.0.1:0")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
port := pc.LocalAddr().(*net.UDPAddr).Port
|
||||||
|
_ = pc.Close()
|
||||||
|
return port
|
||||||
|
}
|
||||||
|
|
||||||
func TestService(t *testing.T) {
|
func TestService(t *testing.T) {
|
||||||
|
lighthousePort := ephemeralUDPPort(t)
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
a := newSimpleService(ca, caKey, "a", netip.MustParseAddr("10.0.0.1"), m{
|
a := newSimpleService(ca, caKey, "a", netip.MustParseAddr("10.0.0.1"), m{
|
||||||
"static_host_map": m{},
|
"static_host_map": m{},
|
||||||
@@ -98,12 +116,12 @@ func TestService(t *testing.T) {
|
|||||||
},
|
},
|
||||||
"listen": m{
|
"listen": m{
|
||||||
"host": "0.0.0.0",
|
"host": "0.0.0.0",
|
||||||
"port": 4243,
|
"port": lighthousePort,
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
b := newSimpleService(ca, caKey, "b", netip.MustParseAddr("10.0.0.2"), m{
|
b := newSimpleService(ca, caKey, "b", netip.MustParseAddr("10.0.0.2"), m{
|
||||||
"static_host_map": m{
|
"static_host_map": m{
|
||||||
"10.0.0.1": []string{"localhost:4243"},
|
"10.0.0.1": []string{fmt.Sprintf("localhost:%d", lighthousePort)},
|
||||||
},
|
},
|
||||||
"lighthouse": m{
|
"lighthouse": m{
|
||||||
"hosts": []string{"10.0.0.1"},
|
"hosts": []string{"10.0.0.1"},
|
||||||
|
|||||||
+38
-2
@@ -8,16 +8,49 @@ import (
|
|||||||
|
|
||||||
const MTU = 9001
|
const MTU = 9001
|
||||||
|
|
||||||
|
// MaxWriteBatch is the largest batch any Conn.WriteBatch implementation is
|
||||||
|
// required to accept. Callers SHOULD NOT pass more than this per call; Linux
|
||||||
|
// backends preallocate sendmmsg scratch sized to this value, so exceeding it
|
||||||
|
// only costs additional sendmmsg chunks within a single WriteBatch call.
|
||||||
|
const MaxWriteBatch = 128
|
||||||
|
|
||||||
|
// RxMeta carries per-packet metadata extracted from the RX path (ancillary
|
||||||
|
// data, kernel offload state, etc.) and passed to EncReader callbacks.
|
||||||
|
// Backends that do not produce a particular signal leave its zero value.
|
||||||
|
//
|
||||||
|
// OuterECN is the 2-bit IP-level ECN codepoint stamped on the carrier
|
||||||
|
// datagram (extracted from IP_TOS / IPV6_TCLASS cmsg on Linux). Zero
|
||||||
|
// means Not-ECT, which is also the value backends without ECN RX support
|
||||||
|
// supply on every packet.
|
||||||
|
type RxMeta struct {
|
||||||
|
OuterECN byte
|
||||||
|
}
|
||||||
|
|
||||||
type EncReader func(
|
type EncReader func(
|
||||||
addr netip.AddrPort,
|
addr netip.AddrPort,
|
||||||
payload []byte,
|
payload []byte,
|
||||||
|
meta RxMeta,
|
||||||
)
|
)
|
||||||
|
|
||||||
type Conn interface {
|
type Conn interface {
|
||||||
Rebind() error
|
Rebind() error
|
||||||
LocalAddr() (netip.AddrPort, error)
|
LocalAddr() (netip.AddrPort, error)
|
||||||
ListenOut(r EncReader) error
|
// ListenOut invokes r for each received packet. On batch-capable
|
||||||
|
// backends (recvmmsg), flush is called after each batch is fully
|
||||||
|
// delivered — callers use it to flush per-batch accumulators such as
|
||||||
|
// TUN write coalescers. Single-packet backends call flush after each
|
||||||
|
// packet. flush must not be nil.
|
||||||
|
ListenOut(r EncReader, flush func()) error
|
||||||
WriteTo(b []byte, addr netip.AddrPort) error
|
WriteTo(b []byte, addr netip.AddrPort) error
|
||||||
|
// WriteBatch sends a contiguous batch of packets, each with its own
|
||||||
|
// destination. bufs and addrs must have the same length. outerECNs may
|
||||||
|
// be nil (treated as all-zero / Not-ECT); when non-nil it must have the
|
||||||
|
// same length as bufs, and outerECNs[i] is the 2-bit IP-level ECN
|
||||||
|
// codepoint to set on packet i's outer header. Linux uses sendmmsg(2)
|
||||||
|
// for a single syscall and attaches the value as IP_TOS / IPV6_TCLASS
|
||||||
|
// cmsg; other backends ignore it. Returns on the first error; callers
|
||||||
|
// may observe a partial send if some packets went out before the error.
|
||||||
|
WriteBatch(bufs [][]byte, addrs []netip.AddrPort, outerECNs []byte) error
|
||||||
ReloadConfig(c *config.C)
|
ReloadConfig(c *config.C)
|
||||||
SupportsMultipleReaders() bool
|
SupportsMultipleReaders() bool
|
||||||
Close() error
|
Close() error
|
||||||
@@ -31,7 +64,7 @@ func (NoopConn) Rebind() error {
|
|||||||
func (NoopConn) LocalAddr() (netip.AddrPort, error) {
|
func (NoopConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
return netip.AddrPort{}, nil
|
return netip.AddrPort{}, nil
|
||||||
}
|
}
|
||||||
func (NoopConn) ListenOut(_ EncReader) error {
|
func (NoopConn) ListenOut(_ EncReader, _ func()) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
func (NoopConn) SupportsMultipleReaders() bool {
|
func (NoopConn) SupportsMultipleReaders() bool {
|
||||||
@@ -40,6 +73,9 @@ func (NoopConn) SupportsMultipleReaders() bool {
|
|||||||
func (NoopConn) WriteTo(_ []byte, _ netip.AddrPort) error {
|
func (NoopConn) WriteTo(_ []byte, _ netip.AddrPort) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
func (NoopConn) WriteBatch(_ [][]byte, _ []netip.AddrPort, _ []byte) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
func (NoopConn) ReloadConfig(_ *config.C) {
|
func (NoopConn) ReloadConfig(_ *config.C) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,62 @@
|
|||||||
|
//go:build !android && !e2e_testing
|
||||||
|
// +build !android,!e2e_testing
|
||||||
|
|
||||||
|
package udp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"syscall"
|
||||||
|
"unsafe"
|
||||||
|
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
// rawSendmmsg performs sendmmsg(2) over a syscall.RawConn without
|
||||||
|
// allocating a closure per call. The struct holds preallocated in/out
|
||||||
|
// scratch (chunk/sent/errno) and a method-value bound at construction so
|
||||||
|
// rawConn.Write receives a stable function pointer instead of a fresh
|
||||||
|
// closure on every send.
|
||||||
|
type rawSendmmsg struct {
|
||||||
|
msgs []rawMessage
|
||||||
|
chunk int
|
||||||
|
sent int
|
||||||
|
errno syscall.Errno
|
||||||
|
callback func(fd uintptr) bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// bind wires r.callback to r.run. Must be called once after r.msgs is set;
|
||||||
|
// subsequent send calls invoke r.callback without rebinding.
|
||||||
|
func (r *rawSendmmsg) bind() { r.callback = r.run }
|
||||||
|
|
||||||
|
// run is the preallocated callback rawConn.Write invokes. It reads its
|
||||||
|
// input (r.chunk) and writes its outputs (r.sent, r.errno) through the
|
||||||
|
// rawSendmmsg fields so the method value does not capture per-call locals
|
||||||
|
// and therefore does not heap-allocate.
|
||||||
|
func (r *rawSendmmsg) run(fd uintptr) bool {
|
||||||
|
r1, _, errno := unix.Syscall6(unix.SYS_SENDMMSG, fd,
|
||||||
|
uintptr(unsafe.Pointer(&r.msgs[0])), uintptr(r.chunk),
|
||||||
|
0, 0, 0,
|
||||||
|
)
|
||||||
|
if errno == syscall.EAGAIN || errno == syscall.EWOULDBLOCK {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
r.sent = int(r1)
|
||||||
|
r.errno = errno
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// send issues sendmmsg over rc against the first n entries of r.msgs.
|
||||||
|
// Returns the number of entries the kernel processed and any error;
|
||||||
|
// matches the original sendmmsg helper's contract.
|
||||||
|
func (r *rawSendmmsg) send(rc syscall.RawConn, n int) (int, error) {
|
||||||
|
r.chunk = n
|
||||||
|
r.sent = 0
|
||||||
|
r.errno = 0
|
||||||
|
if err := rc.Write(r.callback); err != nil {
|
||||||
|
return r.sent, err
|
||||||
|
}
|
||||||
|
if r.errno != 0 {
|
||||||
|
return r.sent, &net.OpError{Op: "sendmmsg", Err: r.errno}
|
||||||
|
}
|
||||||
|
return r.sent, nil
|
||||||
|
}
|
||||||
+12
-2
@@ -140,6 +140,15 @@ func (u *StdConn) WriteTo(b []byte, ap netip.AddrPort) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||||
|
for i, b := range bufs {
|
||||||
|
if err := u.WriteTo(b, addrs[i]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
a := u.UDPConn.LocalAddr()
|
a := u.UDPConn.LocalAddr()
|
||||||
|
|
||||||
@@ -165,7 +174,7 @@ func NewUDPStatsEmitter(udpConns []Conn) func() {
|
|||||||
return func() {}
|
return func() {}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) ListenOut(r EncReader) error {
|
func (u *StdConn) ListenOut(r EncReader, flush func()) error {
|
||||||
buffer := make([]byte, MTU)
|
buffer := make([]byte, MTU)
|
||||||
|
|
||||||
for {
|
for {
|
||||||
@@ -179,7 +188,8 @@ func (u *StdConn) ListenOut(r EncReader) error {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n])
|
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n], RxMeta{})
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,61 @@
|
|||||||
|
//go:build linux && !android && !e2e_testing
|
||||||
|
|
||||||
|
package udp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestPlanRunBreaksOnECNChange confirms that two same-destination, same-size
|
||||||
|
// packets with different outer ECN end up in separate sendmmsg entries (the
|
||||||
|
// kernel stamps one outer codepoint per entry, so a run that straddled the
|
||||||
|
// boundary would silently lose information).
|
||||||
|
func TestPlanRunBreaksOnECNChange(t *testing.T) {
|
||||||
|
u := &StdConn{gsoSupported: true, maxGSOSegments: 63}
|
||||||
|
dst := netip.MustParseAddrPort("10.0.0.1:4242")
|
||||||
|
|
||||||
|
bufs := [][]byte{
|
||||||
|
make([]byte, 1200),
|
||||||
|
make([]byte, 1200),
|
||||||
|
make([]byte, 1200),
|
||||||
|
}
|
||||||
|
addrs := []netip.AddrPort{dst, dst, dst}
|
||||||
|
|
||||||
|
t.Run("uniform_ecn_runs_together", func(t *testing.T) {
|
||||||
|
ecns := []byte{0x02, 0x02, 0x02}
|
||||||
|
runLen, segSize := u.planRun(bufs, addrs, ecns, 0, 64)
|
||||||
|
if runLen != 3 {
|
||||||
|
t.Errorf("runLen=%d want 3 (uniform ECT(0))", runLen)
|
||||||
|
}
|
||||||
|
if segSize != 1200 {
|
||||||
|
t.Errorf("segSize=%d want 1200", segSize)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("ecn_change_truncates_run", func(t *testing.T) {
|
||||||
|
// 0,0,3: first two run together, CE seeds a fresh entry.
|
||||||
|
ecns := []byte{0x00, 0x00, 0x03}
|
||||||
|
runLen, _ := u.planRun(bufs, addrs, ecns, 0, 64)
|
||||||
|
if runLen != 2 {
|
||||||
|
t.Errorf("runLen=%d want 2 (ECN changes at index 2)", runLen)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("nil_ecns_runs_full", func(t *testing.T) {
|
||||||
|
runLen, _ := u.planRun(bufs, addrs, nil, 0, 64)
|
||||||
|
if runLen != 3 {
|
||||||
|
t.Errorf("runLen=%d want 3 (nil ecns means no break)", runLen)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("first_ecn_is_singleton", func(t *testing.T) {
|
||||||
|
// Second packet has different ECN from the first → run halts at 1
|
||||||
|
// (the first packet alone forms the run).
|
||||||
|
ecns := []byte{0x00, 0x03, 0x03}
|
||||||
|
runLen, _ := u.planRun(bufs, addrs, ecns, 0, 64)
|
||||||
|
if runLen != 1 {
|
||||||
|
t.Errorf("runLen=%d want 1 (different ECN immediately)", runLen)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
+12
-2
@@ -44,6 +44,15 @@ func (u *GenericConn) WriteTo(b []byte, addr netip.AddrPort) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (u *GenericConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||||
|
for i, b := range bufs {
|
||||||
|
if _, err := u.UDPConn.WriteToUDPAddrPort(b, addrs[i]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (u *GenericConn) LocalAddr() (netip.AddrPort, error) {
|
func (u *GenericConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
a := u.UDPConn.LocalAddr()
|
a := u.UDPConn.LocalAddr()
|
||||||
|
|
||||||
@@ -73,7 +82,7 @@ type rawMessage struct {
|
|||||||
Len uint32
|
Len uint32
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *GenericConn) ListenOut(r EncReader) error {
|
func (u *GenericConn) ListenOut(r EncReader, flush func()) error {
|
||||||
buffer := make([]byte, MTU)
|
buffer := make([]byte, MTU)
|
||||||
|
|
||||||
var lastRecvErr time.Time
|
var lastRecvErr time.Time
|
||||||
@@ -93,7 +102,8 @@ func (u *GenericConn) ListenOut(r EncReader) error {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n])
|
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n], RxMeta{})
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+698
-17
@@ -6,10 +6,13 @@ package udp
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
"syscall"
|
"syscall"
|
||||||
"unsafe"
|
"unsafe"
|
||||||
|
|
||||||
@@ -24,6 +27,52 @@ type StdConn struct {
|
|||||||
isV4 bool
|
isV4 bool
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
batch int
|
batch int
|
||||||
|
|
||||||
|
// sendmmsg scratch. Each queue has its own StdConn, so no locking is
|
||||||
|
// needed. Sized to MaxWriteBatch at construction; WriteBatch chunks
|
||||||
|
// larger inputs.
|
||||||
|
writeMsgs []rawMessage
|
||||||
|
writeIovs []iovec
|
||||||
|
writeNames [][]byte
|
||||||
|
|
||||||
|
// Per-entry cmsg scratch. writeCmsg is one contiguous slab of
|
||||||
|
// MaxWriteBatch * writeCmsgSpace bytes; each entry holds two cmsg
|
||||||
|
// headers (UDP_SEGMENT then IP_TOS / IPV6_TCLASS) pre-filled once in
|
||||||
|
// prepareWriteMessages. WriteBatch only rewrites the per-call data
|
||||||
|
// payloads and toggles Hdr.Control / Hdr.Controllen to point at
|
||||||
|
// whichever subset of the two cmsgs applies.
|
||||||
|
writeCmsg []byte
|
||||||
|
writeCmsgSpace int
|
||||||
|
writeCmsgSegSpace int
|
||||||
|
writeCmsgEcnSpace int
|
||||||
|
|
||||||
|
// writeEntryEnd[e] is the bufs index *after* the last packet packed
|
||||||
|
// into mmsghdr entry e. Used to rewind `i` on partial sendmmsg success.
|
||||||
|
writeEntryEnd []int
|
||||||
|
|
||||||
|
// rawSend wraps the sendmmsg(2) callback in a closure-free helper so
|
||||||
|
// the hot path doesn't heap-allocate a fresh closure per call.
|
||||||
|
rawSend rawSendmmsg
|
||||||
|
|
||||||
|
// UDP GSO (sendmsg with UDP_SEGMENT cmsg) support. gsoSupported is
|
||||||
|
// probed once at socket creation. When true, WriteBatch packs same-
|
||||||
|
// destination consecutive packets into a single sendmmsg entry with a
|
||||||
|
// UDP_SEGMENT cmsg; otherwise each packet is its own entry.
|
||||||
|
gsoSupported bool
|
||||||
|
maxGSOSegments int
|
||||||
|
|
||||||
|
// UDP GRO (recvmsg with UDP_GRO cmsg) support. groSupported is probed
|
||||||
|
// once at socket creation. When true, listenOutBatch allocates larger
|
||||||
|
// RX buffers and a per-entry cmsg slot so the kernel can coalesce
|
||||||
|
// consecutive same-flow datagrams into a single recvmmsg entry; the
|
||||||
|
// delivered cmsg carries the gso_size used to split them back apart.
|
||||||
|
groSupported bool
|
||||||
|
|
||||||
|
// ecnRecvSupported is true when IP_RECVTOS / IPV6_RECVTCLASS was
|
||||||
|
// successfully enabled — the kernel will deliver the outer IP-ECN of
|
||||||
|
// each arriving datagram as a per-slot cmsg, and listenOutBatch passes
|
||||||
|
// the parsed value to the EncReader callback for RFC 6040 combine.
|
||||||
|
ecnRecvSupported bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func setReusePort(network, address string, c syscall.RawConn) error {
|
func setReusePort(network, address string, c syscall.RawConn) error {
|
||||||
@@ -57,10 +106,11 @@ func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int)
|
|||||||
}
|
}
|
||||||
//gotta find out if we got an AF_INET6 socket or not:
|
//gotta find out if we got an AF_INET6 socket or not:
|
||||||
out := &StdConn{
|
out := &StdConn{
|
||||||
udpConn: udpConn,
|
udpConn: udpConn,
|
||||||
rawConn: rawConn,
|
rawConn: rawConn,
|
||||||
l: l,
|
l: l,
|
||||||
batch: batch,
|
batch: batch,
|
||||||
|
maxGSOSegments: 1,
|
||||||
}
|
}
|
||||||
|
|
||||||
af, err := out.getSockOptInt(unix.SO_DOMAIN)
|
af, err := out.getSockOptInt(unix.SO_DOMAIN)
|
||||||
@@ -70,9 +120,229 @@ func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int)
|
|||||||
}
|
}
|
||||||
out.isV4 = af == unix.AF_INET
|
out.isV4 = af == unix.AF_INET
|
||||||
|
|
||||||
|
out.prepareWriteMessages(MaxWriteBatch)
|
||||||
|
out.rawSend.msgs = out.writeMsgs
|
||||||
|
out.rawSend.bind()
|
||||||
|
|
||||||
|
out.prepareGSO()
|
||||||
|
// GRO delivers coalesced superpackets that need a cmsg to split back
|
||||||
|
// into segments. The single-packet RX path uses ReadFromUDPAddrPort
|
||||||
|
// and cannot see that cmsg, so only enable GRO for the batch path.
|
||||||
|
if batch > 1 {
|
||||||
|
out.prepareGRO()
|
||||||
|
}
|
||||||
|
// Best-effort: ask the kernel to deliver outer IP-ECN as ancillary data
|
||||||
|
// on every recvmmsg slot so the decap side can apply RFC 6040 combine.
|
||||||
|
// On older kernels these may not exist; failing here just means we get
|
||||||
|
// 0 (Not-ECT) on every slot, which is the same as ecn_mode=disable.
|
||||||
|
out.prepareECNRecv()
|
||||||
|
|
||||||
return out, nil
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// prepareWriteMessages allocates one mmsghdr/iovec/sockaddr/cmsg scratch
|
||||||
|
// slot per sendmmsg entry. The iovec slab is sized to n so all entries'
|
||||||
|
// iovecs share one allocation; per-entry fan-out is further capped at
|
||||||
|
// maxGSOSegments. Hdr.Iov / Hdr.Iovlen / Hdr.Control / Hdr.Controllen are
|
||||||
|
// wired per call since each entry can span a variable number of iovecs
|
||||||
|
// and may or may not carry a cmsg.
|
||||||
|
//
|
||||||
|
// Per-mmsghdr cmsg layout. Each entry's slot of length writeCmsgSpace holds
|
||||||
|
// up to two cmsg headers placed at fixed offsets:
|
||||||
|
//
|
||||||
|
// [0 .. writeCmsgSegSpace) UDP_SEGMENT (gso_size, uint16)
|
||||||
|
// [writeCmsgSegSpace .. writeCmsgSpace) IP_TOS or IPV6_TCLASS (int32)
|
||||||
|
//
|
||||||
|
// Both headers are pre-filled once here; per-call we only rewrite the data
|
||||||
|
// payload and toggle Hdr.Control / Hdr.Controllen to point at whichever
|
||||||
|
// subset applies (none / segment-only / ecn-only / both).
|
||||||
|
func (u *StdConn) prepareWriteMessages(n int) {
|
||||||
|
u.writeMsgs = make([]rawMessage, n)
|
||||||
|
u.writeIovs = make([]iovec, n)
|
||||||
|
u.writeNames = make([][]byte, n)
|
||||||
|
u.writeEntryEnd = make([]int, n)
|
||||||
|
|
||||||
|
u.writeCmsgSegSpace = unix.CmsgSpace(2)
|
||||||
|
u.writeCmsgEcnSpace = unix.CmsgSpace(4)
|
||||||
|
u.writeCmsgSpace = u.writeCmsgSegSpace + u.writeCmsgEcnSpace
|
||||||
|
u.writeCmsg = make([]byte, n*u.writeCmsgSpace)
|
||||||
|
|
||||||
|
// Default the ECN header to the socket's own family. writeEntryCmsg
|
||||||
|
// finalizes Level/Type per entry from the destination address (a v4-mapped
|
||||||
|
// dst on a dual-stack v6 socket needs IP_TOS, not IPV6_TCLASS), so this is
|
||||||
|
// only the value used before the first per-entry rewrite.
|
||||||
|
ecnLevel := int32(unix.IPPROTO_IP)
|
||||||
|
ecnType := int32(unix.IP_TOS)
|
||||||
|
if !u.isV4 {
|
||||||
|
ecnLevel = unix.IPPROTO_IPV6
|
||||||
|
ecnType = unix.IPV6_TCLASS
|
||||||
|
}
|
||||||
|
|
||||||
|
for k := 0; k < n; k++ {
|
||||||
|
base := k * u.writeCmsgSpace
|
||||||
|
seg := (*unix.Cmsghdr)(unsafe.Pointer(&u.writeCmsg[base]))
|
||||||
|
seg.Level = unix.SOL_UDP
|
||||||
|
seg.Type = unix.UDP_SEGMENT
|
||||||
|
setCmsgLen(seg, unix.CmsgLen(2))
|
||||||
|
|
||||||
|
ecn := (*unix.Cmsghdr)(unsafe.Pointer(&u.writeCmsg[base+u.writeCmsgSegSpace]))
|
||||||
|
ecn.Level = ecnLevel
|
||||||
|
ecn.Type = ecnType
|
||||||
|
setCmsgLen(ecn, unix.CmsgLen(4))
|
||||||
|
}
|
||||||
|
|
||||||
|
for i := range u.writeMsgs {
|
||||||
|
u.writeNames[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||||
|
u.writeMsgs[i].Hdr.Name = &u.writeNames[i][0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// maxGSOBytes bounds the total payload per sendmsg() when UDP_SEGMENT is
|
||||||
|
// set. The kernel stitches all iovecs into a single skb whose length the
|
||||||
|
// UDP length field can represent, and also enforces sk_gso_max_size (which
|
||||||
|
// on most devices is 65536). We use 65000 to leave headroom under the
|
||||||
|
// 65535 UDP-length cap, avoiding EMSGSIZE on large TSO superpackets.
|
||||||
|
const maxGSOBytes = 65000
|
||||||
|
|
||||||
|
// prepareGSO probes UDP_SEGMENT support and sets u.gsoSupported on success.
|
||||||
|
// Best-effort; failure leaves it false.
|
||||||
|
func (u *StdConn) prepareGSO() {
|
||||||
|
u.maxGSOSegments = 63 //gotta be one less than the max so we can still attach a header
|
||||||
|
|
||||||
|
var probeErr error
|
||||||
|
if err := u.rawConn.Control(func(fd uintptr) {
|
||||||
|
probeErr = unix.SetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_SEGMENT, 0)
|
||||||
|
}); err != nil {
|
||||||
|
u.l.Info("udp: GSO disabled", "reason", "rawconn control failed", "error", err)
|
||||||
|
recordCapability("udp.gso.enabled", false)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if probeErr != nil {
|
||||||
|
u.l.Info("udp: GSO disabled", "reason", "kernel rejected probe", "error", probeErr)
|
||||||
|
recordCapability("udp.gso.enabled", false)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var un unix.Utsname
|
||||||
|
if err := unix.Uname(&un); err != nil {
|
||||||
|
u.l.Info("udp: GSO disabled", "reason", "kernel uname probe failed", "error", err)
|
||||||
|
recordCapability("udp.gso.enabled", false)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
u.maxGSOSegments = gsoMaxSegments(string(un.Release[:]))
|
||||||
|
|
||||||
|
u.gsoSupported = true
|
||||||
|
u.l.Info("udp: GSO enabled", "maxGSOSegments", u.maxGSOSegments)
|
||||||
|
recordCapability("udp.gso.enabled", true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// gsoMaxSegments returns the largest number of UDP_SEGMENT segments a single
|
||||||
|
// sendmsg may carry on the running kernel, reserving one segment for the
|
||||||
|
// header. UDP_MAX_SEGMENTS was 64 until Linux v6.9 (commit 1382e3b6a350,
|
||||||
|
// "udp: change maximum number of UDP segments to 128") raised it to 128;
|
||||||
|
// nothing about this changed in 5.5. On kernels older than 6.9 packing more
|
||||||
|
// than 64 segments gets the sendmsg rejected with EINVAL, so cap at 63 there
|
||||||
|
// and only use 127 from 6.9 on. (Maintainer stance: update your kernel if you
|
||||||
|
// want to go fast — this is a plain version gate, not a runtime probe.)
|
||||||
|
func gsoMaxSegments(release string) int {
|
||||||
|
major, minor := parseRelease(release)
|
||||||
|
if major > 6 || (major == 6 && minor >= 9) {
|
||||||
|
return 127
|
||||||
|
}
|
||||||
|
return 63
|
||||||
|
}
|
||||||
|
|
||||||
|
// udpGROBufferSize sizes the per-entry recvmmsg buffer when UDP_GRO is on.
|
||||||
|
// The kernel stitches a run of same-flow datagrams into a single skb whose
|
||||||
|
// length is bounded by sk_gso_max_size (typically 65535); anything larger
|
||||||
|
// would be MSG_TRUNCed. We use the maximum representable UDP length so a
|
||||||
|
// full superpacket always lands intact.
|
||||||
|
const udpGROBufferSize = 65535
|
||||||
|
|
||||||
|
// udpGROCmsgPayload is the size of the UDP_GRO cmsg data delivered by the
|
||||||
|
// kernel: a single int (gso_size in bytes). See udp_cmsg_recv() in
|
||||||
|
// net/ipv4/udp.c.
|
||||||
|
const udpGROCmsgPayload = 4
|
||||||
|
|
||||||
|
// prepareGRO turns on UDP_GRO so the kernel coalesces consecutive same-flow
|
||||||
|
// datagrams into one recvmmsg entry, with a cmsg carrying the gso_size used
|
||||||
|
// to split them back apart on the application side.
|
||||||
|
func (u *StdConn) prepareGRO() {
|
||||||
|
var probeErr error
|
||||||
|
if err := u.rawConn.Control(func(fd uintptr) {
|
||||||
|
probeErr = unix.SetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_GRO, 1)
|
||||||
|
}); err != nil {
|
||||||
|
u.l.Info("udp: GRO disabled", "reason", "rawconn control failed", "error", err)
|
||||||
|
recordCapability("udp.gro.enabled", false)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if probeErr != nil {
|
||||||
|
u.l.Info("udp: GRO disabled", "reason", "kernel rejected probe", "error", probeErr)
|
||||||
|
recordCapability("udp.gro.enabled", false)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
u.groSupported = true
|
||||||
|
u.l.Info("udp: GRO enabled")
|
||||||
|
recordCapability("udp.gro.enabled", true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// prepareECNRecv turns on IP_RECVTOS / IPV6_RECVTCLASS so the outer IP-ECN
|
||||||
|
// field of each arriving datagram is delivered as ancillary data alongside
|
||||||
|
// the payload. listenOutBatch reads it via parseRecvCmsg and passes the
|
||||||
|
// codepoint through the EncReader for RFC 6040 combine on the decap side.
|
||||||
|
// Best-effort: we keep going on failure.
|
||||||
|
func (u *StdConn) prepareECNRecv() {
|
||||||
|
var v4err, v6err error
|
||||||
|
if err := u.rawConn.Control(func(fd uintptr) {
|
||||||
|
v4err = unix.SetsockoptInt(int(fd), unix.IPPROTO_IP, unix.IP_RECVTOS, 1)
|
||||||
|
if !u.isV4 {
|
||||||
|
v6err = unix.SetsockoptInt(int(fd), unix.IPPROTO_IPV6, unix.IPV6_RECVTCLASS, 1)
|
||||||
|
}
|
||||||
|
}); err != nil {
|
||||||
|
u.l.Info("udp: outer-ECN RX disabled", "reason", "rawconn control failed", "error", err)
|
||||||
|
recordCapability("udp.ecn_rx.enabled", false)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if u.isV4 { //only check the V4 attempt
|
||||||
|
if v4err != nil {
|
||||||
|
u.l.Info("udp: outer-ECN RX disabled", "reason", "kernel rejected probe", "error", v4err)
|
||||||
|
recordCapability("udp.ecn_rx.enabled", false)
|
||||||
|
} else {
|
||||||
|
u.ecnRecvSupported = true
|
||||||
|
u.l.Info("udp: outer-ECN RX enabled")
|
||||||
|
recordCapability("udp.ecn_rx.enabled", true)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
} else {
|
||||||
|
if v6err != nil { //no V6 ECN? disable it.
|
||||||
|
u.l.Info("udp: outer-ECN RX disabled", "reason", "kernel rejected probe", "error", errors.Join(v4err, v6err))
|
||||||
|
recordCapability("udp.ecn_rx.enabled", false)
|
||||||
|
return
|
||||||
|
} else if v4err != nil { //no V4, but yes V6? Low level warning. Could be a V6-specific bind.
|
||||||
|
u.l.Debug("udp: outer-ECN RX degraded", "reason", "kernel rejected probe on IPv4", "error", v4err)
|
||||||
|
}
|
||||||
|
// all good
|
||||||
|
u.ecnRecvSupported = true
|
||||||
|
u.l.Info("udp: outer-ECN RX enabled")
|
||||||
|
recordCapability("udp.ecn_rx.enabled", true)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// recordCapability registers (or updates) a boolean gauge for one of the
|
||||||
|
// kernel-feature probes. Gauges go to 1 when the feature is enabled, 0 when
|
||||||
|
// it is not — dashboards can show degraded state on partially-supported
|
||||||
|
// kernels at a glance. Calling repeatedly with the same name updates the
|
||||||
|
// existing gauge rather than registering a duplicate.
|
||||||
|
func recordCapability(name string, enabled bool) {
|
||||||
|
g := metrics.GetOrRegisterGauge(name, nil)
|
||||||
|
if enabled {
|
||||||
|
g.Update(1)
|
||||||
|
} else {
|
||||||
|
g.Update(0)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (u *StdConn) SupportsMultipleReaders() bool {
|
func (u *StdConn) SupportsMultipleReaders() bool {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
@@ -171,7 +441,7 @@ func recvmmsg(fd uintptr, msgs []rawMessage) (int, bool, error) {
|
|||||||
return int(n), true, nil
|
return int(n), true, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) listenOutSingle(r EncReader) error {
|
func (u *StdConn) listenOutSingle(r EncReader, flush func()) error {
|
||||||
var err error
|
var err error
|
||||||
var n int
|
var n int
|
||||||
var from netip.AddrPort
|
var from netip.AddrPort
|
||||||
@@ -183,16 +453,42 @@ func (u *StdConn) listenOutSingle(r EncReader) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
from = netip.AddrPortFrom(from.Addr().Unmap(), from.Port())
|
from = netip.AddrPortFrom(from.Addr().Unmap(), from.Port())
|
||||||
r(from, buffer[:n])
|
// listenOutSingle uses ReadFromUDPAddrPort which discards cmsgs,
|
||||||
|
// so the outer ECN field is not visible on this path. Zero RxMeta
|
||||||
|
// (Not-ECT) means RFC 6040 combine is a no-op.
|
||||||
|
r(from, buffer[:n], RxMeta{})
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) listenOutBatch(r EncReader) error {
|
func getFrom(names [][]byte, i int, isV4 bool) netip.AddrPort {
|
||||||
var ip netip.Addr
|
var ip netip.Addr
|
||||||
|
// Its ok to skip the ok check here, the slicing is the only error that can occur and it will panic
|
||||||
|
if isV4 {
|
||||||
|
ip, _ = netip.AddrFromSlice(names[i][4:8])
|
||||||
|
} else {
|
||||||
|
ip, _ = netip.AddrFromSlice(names[i][8:24])
|
||||||
|
}
|
||||||
|
return netip.AddrPortFrom(ip.Unmap(), binary.BigEndian.Uint16(names[i][2:4]))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (u *StdConn) listenOutBatch(r EncReader, flush func()) error {
|
||||||
var n int
|
var n int
|
||||||
var operr error
|
var operr error
|
||||||
|
|
||||||
msgs, buffers, names := u.PrepareRawMessages(u.batch)
|
bufSize := MTU
|
||||||
|
cmsgSpace := 0
|
||||||
|
if u.groSupported {
|
||||||
|
bufSize = udpGROBufferSize
|
||||||
|
cmsgSpace = unix.CmsgSpace(udpGROCmsgPayload)
|
||||||
|
}
|
||||||
|
if u.ecnRecvSupported {
|
||||||
|
// IP_TOS arrives as 1 byte; IPV6_TCLASS arrives as a 4-byte int.
|
||||||
|
// Reserve enough for the wider of the two so the same buffer fits
|
||||||
|
// either family alongside any UDP_GRO cmsg.
|
||||||
|
cmsgSpace += unix.CmsgSpace(4)
|
||||||
|
}
|
||||||
|
msgs, buffers, names, _ := u.PrepareRawMessages(u.batch, bufSize, cmsgSpace)
|
||||||
|
|
||||||
//reader needs to capture variables from this function, since it's used as a lambda with rawConn.Read
|
//reader needs to capture variables from this function, since it's used as a lambda with rawConn.Read
|
||||||
//defining it outside the loop so it gets re-used
|
//defining it outside the loop so it gets re-used
|
||||||
@@ -202,6 +498,11 @@ func (u *StdConn) listenOutBatch(r EncReader) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
for {
|
for {
|
||||||
|
if cmsgSpace > 0 {
|
||||||
|
for i := range msgs {
|
||||||
|
setMsgControllen(&msgs[i].Hdr, cmsgSpace)
|
||||||
|
}
|
||||||
|
}
|
||||||
err := u.rawConn.Read(reader)
|
err := u.rawConn.Read(reader)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -211,22 +512,95 @@ func (u *StdConn) listenOutBatch(r EncReader) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
for i := 0; i < n; i++ {
|
for i := 0; i < n; i++ {
|
||||||
// Its ok to skip the ok check here, the slicing is the only error that can occur and it will panic
|
from := getFrom(names, i, u.isV4)
|
||||||
if u.isV4 {
|
payload := buffers[i][:msgs[i].Len]
|
||||||
ip, _ = netip.AddrFromSlice(names[i][4:8])
|
|
||||||
} else {
|
segSize := 0
|
||||||
ip, _ = netip.AddrFromSlice(names[i][8:24])
|
outerECN := byte(0)
|
||||||
|
if cmsgSpace > 0 {
|
||||||
|
segSize, outerECN = parseRecvCmsg(&msgs[i].Hdr, u.groSupported, u.ecnRecvSupported)
|
||||||
|
}
|
||||||
|
|
||||||
|
if segSize <= 0 || segSize >= len(payload) {
|
||||||
|
r(from, payload, RxMeta{OuterECN: outerECN})
|
||||||
|
} else {
|
||||||
|
for off := 0; off < len(payload); off += segSize {
|
||||||
|
end := off + segSize
|
||||||
|
if end > len(payload) {
|
||||||
|
end = len(payload)
|
||||||
|
}
|
||||||
|
seg := payload[off:end]
|
||||||
|
r(from, seg, RxMeta{OuterECN: outerECN})
|
||||||
|
}
|
||||||
}
|
}
|
||||||
r(netip.AddrPortFrom(ip.Unmap(), binary.BigEndian.Uint16(names[i][2:4])), buffers[i][:msgs[i].Len])
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) ListenOut(r EncReader) error {
|
// headerCounter returns the big-endian uint64 message counter at bytes
|
||||||
|
// [8:16] of a nebula packet, or 0 if the buffer is too short.
|
||||||
|
func headerCounter(buf []byte) uint64 {
|
||||||
|
if len(buf) < 16 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return binary.BigEndian.Uint64(buf[8:16])
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseRecvCmsg walks the per-slot ancillary buffer once and extracts up to
|
||||||
|
// two values of interest in a single pass: the UDP_GRO gso_size (when
|
||||||
|
// wantGRO is true) and the outer IP-level ECN codepoint stamped on the
|
||||||
|
// carrier (when wantECN is true). Returns zeros for whichever field is not
|
||||||
|
// requested or not present.
|
||||||
|
//
|
||||||
|
// The outer ECN is accepted from EITHER an IP_TOS (IPPROTO_IP, 1-byte) or an
|
||||||
|
// IPV6_TCLASS (IPPROTO_IPV6, 4-byte int) cmsg, regardless of the socket's
|
||||||
|
// family: a dual-stack v6 socket (isV4 == false) delivers IPv4 peers' outer
|
||||||
|
// ECN as an IP_TOS cmsg — gating on socket family here dropped v4-underlay
|
||||||
|
// ECN entirely. Whichever cmsg the kernel delivered carries the value.
|
||||||
|
func parseRecvCmsg(hdr *msghdr, wantGRO, wantECN bool) (gso int, ecn byte) {
|
||||||
|
controllen := int(hdr.Controllen)
|
||||||
|
if controllen < unix.SizeofCmsghdr || hdr.Control == nil {
|
||||||
|
return 0, 0
|
||||||
|
}
|
||||||
|
ctrl := unsafe.Slice(hdr.Control, controllen)
|
||||||
|
off := 0
|
||||||
|
for off+unix.SizeofCmsghdr <= len(ctrl) {
|
||||||
|
ch := (*unix.Cmsghdr)(unsafe.Pointer(&ctrl[off]))
|
||||||
|
clen := int(ch.Len)
|
||||||
|
if clen < unix.SizeofCmsghdr || off+clen > len(ctrl) {
|
||||||
|
return gso, ecn
|
||||||
|
}
|
||||||
|
dataOff := off + unix.CmsgLen(0)
|
||||||
|
switch {
|
||||||
|
case wantGRO && ch.Level == unix.SOL_UDP && ch.Type == unix.UDP_GRO:
|
||||||
|
if dataOff+udpGROCmsgPayload <= len(ctrl) {
|
||||||
|
gso = int(int32(binary.NativeEndian.Uint32(ctrl[dataOff : dataOff+udpGROCmsgPayload])))
|
||||||
|
}
|
||||||
|
case wantECN && ch.Level == unix.IPPROTO_IP && ch.Type == unix.IP_TOS:
|
||||||
|
// IP_TOS arrives as a single byte; only the low 2 bits are ECN.
|
||||||
|
// A dual-stack v6 socket carries v4 peers' outer ECN here.
|
||||||
|
if dataOff+1 <= len(ctrl) {
|
||||||
|
ecn = ctrl[dataOff] & 0x03
|
||||||
|
}
|
||||||
|
case wantECN && ch.Level == unix.IPPROTO_IPV6 && ch.Type == unix.IPV6_TCLASS:
|
||||||
|
// IPV6_TCLASS arrives as a 4-byte int; ECN is the low 2 bits.
|
||||||
|
if dataOff+4 <= len(ctrl) {
|
||||||
|
ecn = byte(binary.NativeEndian.Uint32(ctrl[dataOff:dataOff+4])) & 0x03
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Advance by the aligned cmsg space.
|
||||||
|
off += unix.CmsgSpace(clen - unix.CmsgLen(0))
|
||||||
|
}
|
||||||
|
return gso, ecn
|
||||||
|
}
|
||||||
|
|
||||||
|
func (u *StdConn) ListenOut(r EncReader, flush func()) error {
|
||||||
if u.batch == 1 {
|
if u.batch == 1 {
|
||||||
return u.listenOutSingle(r)
|
return u.listenOutSingle(r, flush)
|
||||||
} else {
|
} else {
|
||||||
return u.listenOutBatch(r)
|
return u.listenOutBatch(r, flush)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -235,6 +609,294 @@ func (u *StdConn) WriteTo(b []byte, ip netip.AddrPort) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// WriteBatch sends bufs via sendmmsg(2) using the preallocated scratch on
|
||||||
|
// StdConn. Consecutive packets to the same destination with matching segment
|
||||||
|
// sizes (all but possibly the last) are coalesced into a single mmsghdr entry
|
||||||
|
// carrying a UDP_SEGMENT cmsg, so one syscall can mix runs of GSO superpackets
|
||||||
|
// with plain one-off datagrams. Without GSO support every packet is its own
|
||||||
|
// entry, matching the prior behaviour.
|
||||||
|
//
|
||||||
|
// Chunks larger than the scratch are processed across multiple syscalls. If
|
||||||
|
// sendmmsg returns an error AND zero entries went out we fall back to
|
||||||
|
// per-packet WriteTo for that chunk so the caller still gets best-effort
|
||||||
|
// delivery; on a partial-success error we just replay the remainder.
|
||||||
|
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, ecns []byte) error {
|
||||||
|
if len(bufs) != len(addrs) {
|
||||||
|
return fmt.Errorf("WriteBatch: len(bufs)=%d != len(addrs)=%d", len(bufs), len(addrs))
|
||||||
|
}
|
||||||
|
if ecns != nil && len(ecns) != len(bufs) {
|
||||||
|
return fmt.Errorf("WriteBatch: len(ecns)=%d != len(bufs)=%d", len(ecns), len(bufs))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Callers deliver same-destination packets contiguously and in counter
|
||||||
|
// order, so we run the GSO planner directly without a pre-sort. A
|
||||||
|
// sorting pass measurably hurt throughput in microbenchmarks while
|
||||||
|
// providing no observed reordering benefit.
|
||||||
|
|
||||||
|
i := 0
|
||||||
|
sendChunks:
|
||||||
|
for i < len(bufs) {
|
||||||
|
baseI := i
|
||||||
|
entry := 0
|
||||||
|
iovIdx := 0
|
||||||
|
for entry < len(u.writeMsgs) && i < len(bufs) {
|
||||||
|
iovBudget := len(u.writeIovs) - iovIdx
|
||||||
|
if iovBudget < 1 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
runLen, segSize := u.planRun(bufs, addrs, ecns, i, iovBudget)
|
||||||
|
if runLen == 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
for k := 0; k < runLen; k++ {
|
||||||
|
b := bufs[i+k]
|
||||||
|
if len(b) == 0 {
|
||||||
|
u.writeIovs[iovIdx+k].Base = nil
|
||||||
|
setIovLen(&u.writeIovs[iovIdx+k], 0)
|
||||||
|
} else {
|
||||||
|
u.writeIovs[iovIdx+k].Base = &b[0]
|
||||||
|
setIovLen(&u.writeIovs[iovIdx+k], len(b))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
nlen, err := writeSockaddr(u.writeNames[entry], addrs[i], u.isV4)
|
||||||
|
if err != nil {
|
||||||
|
// One destination in this chunk has an address family the
|
||||||
|
// socket can't send to (e.g. an IPv6 remote on a v4-bound
|
||||||
|
// socket → ErrInvalidIPv6RemoteForSocket). Abandoning the whole
|
||||||
|
// sendmmsg here would drop every packet already packed for this
|
||||||
|
// chunk plus every packet still ahead of us in bufs. Instead
|
||||||
|
// fall back to per-packet WriteTo for the packets packed so far
|
||||||
|
// in this chunk and the offending one: WriteTo delivers each
|
||||||
|
// good destination and only errors on the bad one, which we
|
||||||
|
// drop and keep going. One bad destination costs one packet,
|
||||||
|
// never the batch. (Same fallback the zero-sent sendmmsg path
|
||||||
|
// below uses, extended to cover the misaddressed packet.)
|
||||||
|
for k := baseI; k <= i; k++ {
|
||||||
|
if werr := u.WriteTo(bufs[k], addrs[k]); werr != nil && k != i {
|
||||||
|
return werr
|
||||||
|
}
|
||||||
|
}
|
||||||
|
i++
|
||||||
|
continue sendChunks
|
||||||
|
}
|
||||||
|
|
||||||
|
hdr := &u.writeMsgs[entry].Hdr
|
||||||
|
hdr.Iov = &u.writeIovs[iovIdx]
|
||||||
|
setMsgIovlen(hdr, runLen)
|
||||||
|
hdr.Namelen = uint32(nlen)
|
||||||
|
|
||||||
|
var ecn byte
|
||||||
|
if ecns != nil {
|
||||||
|
ecn = ecns[i]
|
||||||
|
}
|
||||||
|
// ECN cmsg family follows the destination, not the socket: a
|
||||||
|
// v4-mapped dst on a dual-stack v6 socket must be stamped via
|
||||||
|
// IP_TOS. addrs[i] is this run's destination (i advances below).
|
||||||
|
dstIsV4 := addrs[i].Addr().Unmap().Is4()
|
||||||
|
u.writeEntryCmsg(entry, runLen, segSize, ecn, dstIsV4)
|
||||||
|
|
||||||
|
i += runLen
|
||||||
|
iovIdx += runLen
|
||||||
|
u.writeEntryEnd[entry] = i
|
||||||
|
entry++
|
||||||
|
}
|
||||||
|
|
||||||
|
if entry == 0 {
|
||||||
|
return fmt.Errorf("sendmmsg: no progress")
|
||||||
|
}
|
||||||
|
|
||||||
|
sent, serr := u.sendmmsg(entry)
|
||||||
|
if serr != nil && sent <= 0 {
|
||||||
|
// Nothing went out for this chunk; fall back to WriteTo for each
|
||||||
|
// packet that was queued this iteration. We only enter this path
|
||||||
|
// when sendmmsg returned an error AND zero entries succeeded —
|
||||||
|
// otherwise the partial-success advance below replays only the
|
||||||
|
// remainder, avoiding duplicates of already-sent packets.
|
||||||
|
//
|
||||||
|
// sent=-1 from sendmmsg means message 0 itself failed (partial
|
||||||
|
// success returns the count instead), so log entry 0's parameters
|
||||||
|
// — that's the entry the kernel rejected.
|
||||||
|
hdr0 := &u.writeMsgs[0].Hdr
|
||||||
|
runLen0 := u.writeEntryEnd[0] - baseI
|
||||||
|
seg0 := len(bufs[baseI])
|
||||||
|
ecn0 := byte(0)
|
||||||
|
if ecns != nil {
|
||||||
|
ecn0 = ecns[baseI]
|
||||||
|
}
|
||||||
|
u.l.Warn("sendmmsg had problem",
|
||||||
|
"sent", sent, "err", serr,
|
||||||
|
"entries", entry,
|
||||||
|
"entry0_runLen", runLen0,
|
||||||
|
"entry0_segSize", seg0,
|
||||||
|
"entry0_iovlen", hdr0.Iovlen,
|
||||||
|
"entry0_controllen", hdr0.Controllen,
|
||||||
|
"entry0_namelen", hdr0.Namelen,
|
||||||
|
"entry0_ecn", ecn0,
|
||||||
|
"entry0_dst", addrs[baseI],
|
||||||
|
"isV4", u.isV4,
|
||||||
|
"gso", u.gsoSupported,
|
||||||
|
"gro", u.groSupported,
|
||||||
|
)
|
||||||
|
for k := baseI; k < i; k++ {
|
||||||
|
if werr := u.WriteTo(bufs[k], addrs[k]); werr != nil {
|
||||||
|
return werr
|
||||||
|
}
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if sent == 0 {
|
||||||
|
return fmt.Errorf("sendmmsg made no progress")
|
||||||
|
}
|
||||||
|
// Rewind i to the end of the last successfully sent entry. For a
|
||||||
|
// full-success send this leaves i unchanged; for a partial send it
|
||||||
|
// replays the remainder on the next outer-loop iteration.
|
||||||
|
i = u.writeEntryEnd[sent-1]
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// planRun groups consecutive packets starting at `start` that can be sent as
|
||||||
|
// a single UDP GSO superpacket (one sendmmsg entry with UDP_SEGMENT cmsg).
|
||||||
|
// A run of length 1 means the entry carries no UDP_SEGMENT cmsg and the
|
||||||
|
// kernel treats it as a plain datagram. Returns the run length and the
|
||||||
|
// per-segment size (which equals len(bufs[start])). Without GSO support
|
||||||
|
// every call returns runLen=1. Outer ECN (when ecns != nil) is also a run
|
||||||
|
// boundary — the kernel stamps one outer codepoint per sendmsg entry, so
|
||||||
|
// mixing values inside a run would lose information.
|
||||||
|
func (u *StdConn) planRun(bufs [][]byte, addrs []netip.AddrPort, ecns []byte, start, iovBudget int) (int, int) {
|
||||||
|
if start >= len(bufs) || iovBudget < 1 {
|
||||||
|
return 0, 0
|
||||||
|
}
|
||||||
|
segSize := len(bufs[start])
|
||||||
|
if !u.gsoSupported || segSize == 0 || segSize > maxGSOBytes {
|
||||||
|
return 1, segSize
|
||||||
|
}
|
||||||
|
dst := addrs[start]
|
||||||
|
var ecn byte
|
||||||
|
if ecns != nil {
|
||||||
|
ecn = ecns[start]
|
||||||
|
}
|
||||||
|
maxLen := u.maxGSOSegments
|
||||||
|
if iovBudget < maxLen {
|
||||||
|
maxLen = iovBudget
|
||||||
|
}
|
||||||
|
runLen := 1
|
||||||
|
total := segSize
|
||||||
|
for runLen < maxLen && start+runLen < len(bufs) {
|
||||||
|
nextLen := len(bufs[start+runLen])
|
||||||
|
if nextLen == 0 || nextLen > segSize {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if addrs[start+runLen] != dst {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if ecns != nil && ecns[start+runLen] != ecn {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if total+nextLen > maxGSOBytes {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
total += nextLen
|
||||||
|
runLen++
|
||||||
|
if nextLen < segSize {
|
||||||
|
// A short packet must be the last in the run.
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return runLen, segSize
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeEntryCmsg sets up the per-mmsghdr Hdr.Control / Hdr.Controllen for one
|
||||||
|
// entry. It writes the UDP_SEGMENT payload when runLen >= 2 and the
|
||||||
|
// IP_TOS/IPV6_TCLASS payload when ecn != 0, then points hdr.Control at the
|
||||||
|
// smallest contiguous span that covers whichever cmsg(s) actually apply.
|
||||||
|
//
|
||||||
|
// The outer-ECN cmsg family must match the *destination*, not the socket: on
|
||||||
|
// the default dual-stack v6 bind, a v4-mapped destination is routed through
|
||||||
|
// the kernel's IPv4 path, which parses IP_TOS (IPPROTO_IP) and ignores an
|
||||||
|
// IPV6_TCLASS cmsg. prepareWriteMessages pre-fills a default header; here we
|
||||||
|
// rewrite its Level/Type (and Len) per entry from dstIsV4 so v4 peers get
|
||||||
|
// IP_TOS and v6 peers get IPV6_TCLASS. The data payload is a 4-byte int for
|
||||||
|
// both families, so the pre-computed cmsg space is unchanged.
|
||||||
|
func (u *StdConn) writeEntryCmsg(entry, runLen, segSize int, ecn byte, dstIsV4 bool) {
|
||||||
|
hdr := &u.writeMsgs[entry].Hdr
|
||||||
|
useSeg := runLen >= 2
|
||||||
|
useEcn := ecn != 0
|
||||||
|
base := entry * u.writeCmsgSpace
|
||||||
|
|
||||||
|
if useSeg {
|
||||||
|
dataOff := base + unix.CmsgLen(0)
|
||||||
|
binary.NativeEndian.PutUint16(u.writeCmsg[dataOff:dataOff+2], uint16(segSize))
|
||||||
|
}
|
||||||
|
if useEcn {
|
||||||
|
ecnHdr := (*unix.Cmsghdr)(unsafe.Pointer(&u.writeCmsg[base+u.writeCmsgSegSpace]))
|
||||||
|
if dstIsV4 {
|
||||||
|
ecnHdr.Level = int32(unix.IPPROTO_IP)
|
||||||
|
ecnHdr.Type = int32(unix.IP_TOS)
|
||||||
|
} else {
|
||||||
|
ecnHdr.Level = int32(unix.IPPROTO_IPV6)
|
||||||
|
ecnHdr.Type = int32(unix.IPV6_TCLASS)
|
||||||
|
}
|
||||||
|
setCmsgLen(ecnHdr, unix.CmsgLen(4))
|
||||||
|
dataOff := base + u.writeCmsgSegSpace + unix.CmsgLen(0)
|
||||||
|
binary.NativeEndian.PutUint32(u.writeCmsg[dataOff:dataOff+4], uint32(ecn))
|
||||||
|
}
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case useSeg && useEcn:
|
||||||
|
hdr.Control = &u.writeCmsg[base]
|
||||||
|
setMsgControllen(hdr, u.writeCmsgSpace)
|
||||||
|
case useSeg:
|
||||||
|
hdr.Control = &u.writeCmsg[base]
|
||||||
|
setMsgControllen(hdr, u.writeCmsgSegSpace)
|
||||||
|
case useEcn:
|
||||||
|
hdr.Control = &u.writeCmsg[base+u.writeCmsgSegSpace]
|
||||||
|
setMsgControllen(hdr, u.writeCmsgEcnSpace)
|
||||||
|
default:
|
||||||
|
hdr.Control = nil
|
||||||
|
setMsgControllen(hdr, 0)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// sendmmsg issues sendmmsg(2) over u.rawConn against the first n entries
|
||||||
|
// of u.writeMsgs. Routes through u.rawSend so the per-call kernel callback
|
||||||
|
// stays alloc-free.
|
||||||
|
func (u *StdConn) sendmmsg(n int) (int, error) {
|
||||||
|
return u.rawSend.send(u.rawConn, n)
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeSockaddr encodes addr into buf (which must be at least
|
||||||
|
// SizeofSockaddrInet6 bytes). Returns the number of bytes used. If isV4 is
|
||||||
|
// true and addr is not a v4 (or v4-in-v6) address, returns an error.
|
||||||
|
func writeSockaddr(buf []byte, addr netip.AddrPort, isV4 bool) (int, error) {
|
||||||
|
ap := addr.Addr().Unmap()
|
||||||
|
if isV4 {
|
||||||
|
if !ap.Is4() {
|
||||||
|
return 0, ErrInvalidIPv6RemoteForSocket
|
||||||
|
}
|
||||||
|
// struct sockaddr_in: { sa_family_t(2), in_port_t(2, BE), in_addr(4), zero(8) }
|
||||||
|
// sa_family is host endian.
|
||||||
|
binary.NativeEndian.PutUint16(buf[0:2], unix.AF_INET)
|
||||||
|
binary.BigEndian.PutUint16(buf[2:4], addr.Port())
|
||||||
|
ip4 := ap.As4()
|
||||||
|
copy(buf[4:8], ip4[:])
|
||||||
|
for j := 8; j < 16; j++ {
|
||||||
|
buf[j] = 0
|
||||||
|
}
|
||||||
|
return unix.SizeofSockaddrInet4, nil
|
||||||
|
}
|
||||||
|
// struct sockaddr_in6: { sa_family_t(2), in_port_t(2, BE), flowinfo(4), in6_addr(16), scope_id(4) }
|
||||||
|
binary.NativeEndian.PutUint16(buf[0:2], unix.AF_INET6)
|
||||||
|
binary.BigEndian.PutUint16(buf[2:4], addr.Port())
|
||||||
|
binary.NativeEndian.PutUint32(buf[4:8], 0)
|
||||||
|
ip6 := addr.Addr().As16()
|
||||||
|
copy(buf[8:24], ip6[:])
|
||||||
|
binary.NativeEndian.PutUint32(buf[24:28], 0)
|
||||||
|
return unix.SizeofSockaddrInet6, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (u *StdConn) ReloadConfig(c *config.C) {
|
func (u *StdConn) ReloadConfig(c *config.C) {
|
||||||
b := c.GetInt("listen.read_buffer", 0)
|
b := c.GetInt("listen.read_buffer", 0)
|
||||||
if b > 0 {
|
if b > 0 {
|
||||||
@@ -340,3 +1002,22 @@ func NewUDPStatsEmitter(udpConns []Conn) func() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func parseRelease(r string) (major, minor int) {
|
||||||
|
// strip anything after the second dot or any non-digit
|
||||||
|
parts := strings.SplitN(r, ".", 3)
|
||||||
|
if len(parts) < 2 {
|
||||||
|
return 0, 0
|
||||||
|
}
|
||||||
|
major, _ = strconv.Atoi(parts[0])
|
||||||
|
// minor may have trailing junk like "15-generic"
|
||||||
|
mp := parts[1]
|
||||||
|
for i, c := range mp {
|
||||||
|
if c < '0' || c > '9' {
|
||||||
|
mp = mp[:i]
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
minor, _ = strconv.Atoi(mp)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|||||||
+29
-3
@@ -30,13 +30,18 @@ type rawMessage struct {
|
|||||||
Len uint32
|
Len uint32
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) {
|
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
|
||||||
msgs := make([]rawMessage, n)
|
msgs := make([]rawMessage, n)
|
||||||
buffers := make([][]byte, n)
|
buffers := make([][]byte, n)
|
||||||
names := make([][]byte, n)
|
names := make([][]byte, n)
|
||||||
|
|
||||||
|
var cmsgs []byte
|
||||||
|
if cmsgSpace > 0 {
|
||||||
|
cmsgs = make([]byte, n*cmsgSpace)
|
||||||
|
}
|
||||||
|
|
||||||
for i := range msgs {
|
for i := range msgs {
|
||||||
buffers[i] = make([]byte, MTU)
|
buffers[i] = make([]byte, bufSize)
|
||||||
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||||
|
|
||||||
vs := []iovec{
|
vs := []iovec{
|
||||||
@@ -48,7 +53,28 @@ func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) {
|
|||||||
|
|
||||||
msgs[i].Hdr.Name = &names[i][0]
|
msgs[i].Hdr.Name = &names[i][0]
|
||||||
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
||||||
|
|
||||||
|
if cmsgSpace > 0 {
|
||||||
|
msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace]
|
||||||
|
msgs[i].Hdr.Controllen = uint32(cmsgSpace)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return msgs, buffers, names
|
return msgs, buffers, names, cmsgs
|
||||||
|
}
|
||||||
|
|
||||||
|
func setIovLen(v *iovec, n int) {
|
||||||
|
v.Len = uint32(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setMsgIovlen(m *msghdr, n int) {
|
||||||
|
m.Iovlen = uint32(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setMsgControllen(m *msghdr, n int) {
|
||||||
|
m.Controllen = uint32(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setCmsgLen(h *unix.Cmsghdr, n int) {
|
||||||
|
h.Len = uint32(n)
|
||||||
}
|
}
|
||||||
|
|||||||
+29
-3
@@ -33,13 +33,18 @@ type rawMessage struct {
|
|||||||
Pad0 [4]byte
|
Pad0 [4]byte
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) {
|
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
|
||||||
msgs := make([]rawMessage, n)
|
msgs := make([]rawMessage, n)
|
||||||
buffers := make([][]byte, n)
|
buffers := make([][]byte, n)
|
||||||
names := make([][]byte, n)
|
names := make([][]byte, n)
|
||||||
|
|
||||||
|
var cmsgs []byte
|
||||||
|
if cmsgSpace > 0 {
|
||||||
|
cmsgs = make([]byte, n*cmsgSpace)
|
||||||
|
}
|
||||||
|
|
||||||
for i := range msgs {
|
for i := range msgs {
|
||||||
buffers[i] = make([]byte, MTU)
|
buffers[i] = make([]byte, bufSize)
|
||||||
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||||
|
|
||||||
vs := []iovec{
|
vs := []iovec{
|
||||||
@@ -51,7 +56,28 @@ func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) {
|
|||||||
|
|
||||||
msgs[i].Hdr.Name = &names[i][0]
|
msgs[i].Hdr.Name = &names[i][0]
|
||||||
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
||||||
|
|
||||||
|
if cmsgSpace > 0 {
|
||||||
|
msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace]
|
||||||
|
msgs[i].Hdr.Controllen = uint64(cmsgSpace)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return msgs, buffers, names
|
return msgs, buffers, names, cmsgs
|
||||||
|
}
|
||||||
|
|
||||||
|
func setIovLen(v *iovec, n int) {
|
||||||
|
v.Len = uint64(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setMsgIovlen(m *msghdr, n int) {
|
||||||
|
m.Iovlen = uint64(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setMsgControllen(m *msghdr, n int) {
|
||||||
|
m.Controllen = uint64(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setCmsgLen(h *unix.Cmsghdr, n int) {
|
||||||
|
h.Len = uint64(n)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,233 @@
|
|||||||
|
//go:build linux && !android && !e2e_testing
|
||||||
|
|
||||||
|
package udp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
"syscall"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
"unsafe"
|
||||||
|
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestGSOMaxSegmentsKernelGate pins the corrected kernel-version gate: the
|
||||||
|
// 128-segment cap (127 usable) only lands in Linux v6.9 (commit 1382e3b6a350),
|
||||||
|
// not 5.5. Everything older stays at the conservative 63.
|
||||||
|
func TestGSOMaxSegmentsKernelGate(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
release string
|
||||||
|
want int
|
||||||
|
}{
|
||||||
|
{"5.4.0", 63},
|
||||||
|
{"5.5.0-generic", 63}, // the old bug bumped here — it must not now
|
||||||
|
{"5.15.0", 63},
|
||||||
|
{"6.1.0", 63},
|
||||||
|
{"6.8.0-generic", 63},
|
||||||
|
{"6.9.0", 127},
|
||||||
|
{"6.10.1-arch1-1", 127},
|
||||||
|
{"7.0.5-arch1-1", 127},
|
||||||
|
{"garbage", 63},
|
||||||
|
{"", 63},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
if got := gsoMaxSegments(c.release); got != c.want {
|
||||||
|
t.Errorf("gsoMaxSegments(%q) = %d, want %d", c.release, got, c.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildCmsg lays out a single ancillary cmsg (header + data) in a fresh buffer
|
||||||
|
// the way the kernel would deliver it, so parseRecvCmsg can be exercised
|
||||||
|
// without a live socket.
|
||||||
|
func buildCmsg(level, typ int32, data []byte) []byte {
|
||||||
|
buf := make([]byte, unix.CmsgSpace(len(data)))
|
||||||
|
h := (*unix.Cmsghdr)(unsafe.Pointer(&buf[0]))
|
||||||
|
h.Level = level
|
||||||
|
h.Type = typ
|
||||||
|
setCmsgLen(h, unix.CmsgLen(len(data)))
|
||||||
|
copy(buf[unix.CmsgLen(0):], data)
|
||||||
|
return buf
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestParseRecvCmsgOuterECNFamily is the RX half of the dual-stack ECN fix:
|
||||||
|
// parseRecvCmsg must read the outer ECN from whichever family the kernel
|
||||||
|
// delivered, not from the socket family. On the default `::` dual-stack bind
|
||||||
|
// a v4 peer's outer ECN arrives as an IP_TOS cmsg, which the old socket-family
|
||||||
|
// gate ignored entirely.
|
||||||
|
func TestParseRecvCmsgOuterECNFamily(t *testing.T) {
|
||||||
|
tc := make([]byte, 4)
|
||||||
|
binary.NativeEndian.PutUint32(tc, 0x02)
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
ctrl []byte
|
||||||
|
want byte
|
||||||
|
}{
|
||||||
|
{"ip_tos_ce", buildCmsg(int32(unix.IPPROTO_IP), int32(unix.IP_TOS), []byte{0x03}), 0x03},
|
||||||
|
{"ip_tos_ect0", buildCmsg(int32(unix.IPPROTO_IP), int32(unix.IP_TOS), []byte{0x02}), 0x02},
|
||||||
|
{"ipv6_tclass_ect0", buildCmsg(int32(unix.IPPROTO_IPV6), int32(unix.IPV6_TCLASS), tc), 0x02},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
t.Run(c.name, func(t *testing.T) {
|
||||||
|
hdr := &msghdr{Control: &c.ctrl[0]}
|
||||||
|
setMsgControllen(hdr, len(c.ctrl))
|
||||||
|
gso, ecn := parseRecvCmsg(hdr, false, true)
|
||||||
|
if gso != 0 {
|
||||||
|
t.Errorf("gso = %d, want 0 (no UDP_GRO cmsg present)", gso)
|
||||||
|
}
|
||||||
|
if ecn != c.want {
|
||||||
|
t.Errorf("ecn = 0x%02x, want 0x%02x", ecn, c.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func testLogger() *slog.Logger {
|
||||||
|
return slog.New(slog.DiscardHandler)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWriteBatchBadFamilyDeliversOthers is the H3 regression: a batch that
|
||||||
|
// contains one destination the socket can't reach (an IPv6 remote on a
|
||||||
|
// v4-bound socket) must still deliver every other packet. Before the fix the
|
||||||
|
// writeSockaddr error returned early and dropped the whole chunk.
|
||||||
|
func TestWriteBatchBadFamilyDeliversOthers(t *testing.T) {
|
||||||
|
rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
|
||||||
|
if err != nil {
|
||||||
|
t.Skipf("cannot open v4 receiver (sandbox?): %v", err)
|
||||||
|
}
|
||||||
|
defer rx.Close()
|
||||||
|
rxPort := rx.LocalAddr().(*net.UDPAddr).Port
|
||||||
|
|
||||||
|
// Bind a *non-wildcard* v4 address so Go gives us a genuine AF_INET
|
||||||
|
// socket. A wildcard v4 bind (0.0.0.0) via network "udp" comes up as a
|
||||||
|
// dual-stack AF_INET6 socket on Linux, for which a v6 dest is not a bad
|
||||||
|
// family — which would defeat the point of this test.
|
||||||
|
c, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, false, 1)
|
||||||
|
if err != nil {
|
||||||
|
t.Skipf("cannot open v4 sender (sandbox?): %v", err)
|
||||||
|
}
|
||||||
|
defer c.Close()
|
||||||
|
sender := c.(*StdConn)
|
||||||
|
if !sender.isV4 {
|
||||||
|
t.Fatalf("expected a v4-bound sender socket, got isV4=false")
|
||||||
|
}
|
||||||
|
|
||||||
|
good := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), uint16(rxPort))
|
||||||
|
bad := netip.MustParseAddrPort("[2001:db8::1]:9999") // genuine v6, unreachable on v4 socket
|
||||||
|
|
||||||
|
bufs := [][]byte{[]byte("AAA"), []byte("BBB"), []byte("CCC")}
|
||||||
|
addrs := []netip.AddrPort{good, bad, good}
|
||||||
|
|
||||||
|
if err := sender.WriteBatch(bufs, addrs, nil); err != nil {
|
||||||
|
t.Fatalf("WriteBatch returned error, want nil (bad dest should be isolated): %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
got := map[string]bool{}
|
||||||
|
rx.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||||
|
buf := make([]byte, 64)
|
||||||
|
for i := 0; i < 2; i++ {
|
||||||
|
n, _, rerr := rx.ReadFromUDPAddrPort(buf)
|
||||||
|
if rerr != nil {
|
||||||
|
t.Fatalf("expected 2 delivered packets, read #%d failed: %v", i+1, rerr)
|
||||||
|
}
|
||||||
|
got[string(buf[:n])] = true
|
||||||
|
}
|
||||||
|
if !got["AAA"] || !got["CCC"] {
|
||||||
|
t.Errorf("delivered set = %v, want AAA and CCC both present", got)
|
||||||
|
}
|
||||||
|
if got["BBB"] {
|
||||||
|
t.Errorf("the bad-family packet BBB was somehow delivered")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWriteBatchOuterTOSToV4Mapped is the TX half of the dual-stack ECN fix,
|
||||||
|
// verified against a live kernel: WriteBatch on the default `::` dual-stack
|
||||||
|
// socket, sending to a v4-mapped destination, must stamp the outer ECN via an
|
||||||
|
// IP_TOS cmsg (not IPV6_TCLASS, which the kernel's v4 path ignores) so a v4
|
||||||
|
// receiver actually sees it.
|
||||||
|
func TestWriteBatchOuterTOSToV4Mapped(t *testing.T) {
|
||||||
|
rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
|
||||||
|
if err != nil {
|
||||||
|
t.Skipf("cannot open v4 receiver (sandbox?): %v", err)
|
||||||
|
}
|
||||||
|
defer rx.Close()
|
||||||
|
rxPort := rx.LocalAddr().(*net.UDPAddr).Port
|
||||||
|
|
||||||
|
// Ask the kernel to deliver the received outer TOS as ancillary data.
|
||||||
|
rxRaw, err := rx.SyscallConn()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SyscallConn: %v", err)
|
||||||
|
}
|
||||||
|
var soErr error
|
||||||
|
if err := rxRaw.Control(func(fd uintptr) {
|
||||||
|
soErr = unix.SetsockoptInt(int(fd), unix.IPPROTO_IP, unix.IP_RECVTOS, 1)
|
||||||
|
}); err != nil || soErr != nil {
|
||||||
|
t.Skipf("cannot enable IP_RECVTOS (sandbox/kernel?): ctrl=%v so=%v", err, soErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
c, err := NewListener(testLogger(), netip.IPv6Unspecified(), 0, false, 1)
|
||||||
|
if err != nil {
|
||||||
|
t.Skipf("cannot open dual-stack sender (sandbox?): %v", err)
|
||||||
|
}
|
||||||
|
defer c.Close()
|
||||||
|
sender := c.(*StdConn)
|
||||||
|
if sender.isV4 {
|
||||||
|
t.Skipf("sender came up v4-only; need a dual-stack v6 socket for this test")
|
||||||
|
}
|
||||||
|
|
||||||
|
// v4-mapped-in-v6 destination: routed through the kernel's IPv4 path.
|
||||||
|
dst := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), uint16(rxPort))
|
||||||
|
const wantECN = byte(0x02) // ECT(0)
|
||||||
|
|
||||||
|
if err := sender.WriteBatch([][]byte{[]byte("tos-probe")}, []netip.AddrPort{dst}, []byte{wantECN}); err != nil {
|
||||||
|
t.Fatalf("WriteBatch: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read the datagram plus its ancillary TOS.
|
||||||
|
rx.SetReadDeadline(time.Now().Add(3 * time.Second))
|
||||||
|
payload := make([]byte, 128)
|
||||||
|
oob := make([]byte, 512)
|
||||||
|
var n, oobn int
|
||||||
|
var rerr error
|
||||||
|
if err := rxRaw.Read(func(fd uintptr) bool {
|
||||||
|
n, oobn, _, _, rerr = unix.Recvmsg(int(fd), payload, oob, 0)
|
||||||
|
if rerr == syscall.EAGAIN || rerr == syscall.EWOULDBLOCK {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("waiting for datagram failed (no delivery?): %v", err)
|
||||||
|
}
|
||||||
|
if rerr != nil {
|
||||||
|
t.Fatalf("Recvmsg: %v", rerr)
|
||||||
|
}
|
||||||
|
if string(payload[:n]) != "tos-probe" {
|
||||||
|
t.Fatalf("payload = %q, want %q", string(payload[:n]), "tos-probe")
|
||||||
|
}
|
||||||
|
|
||||||
|
cmsgs, err := unix.ParseSocketControlMessage(oob[:oobn])
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ParseSocketControlMessage: %v", err)
|
||||||
|
}
|
||||||
|
found := false
|
||||||
|
var gotTOS byte
|
||||||
|
for _, m := range cmsgs {
|
||||||
|
if m.Header.Level == unix.IPPROTO_IP && m.Header.Type == unix.IP_TOS && len(m.Data) >= 1 {
|
||||||
|
found = true
|
||||||
|
gotTOS = m.Data[0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
t.Fatalf("no IP_TOS cmsg delivered to v4 receiver — outer ECN did not land (%d cmsgs)", len(cmsgs))
|
||||||
|
}
|
||||||
|
if gotTOS&0x03 != wantECN {
|
||||||
|
t.Errorf("received outer TOS = 0x%02x, want low-2-bits = 0x%02x", gotTOS, wantECN)
|
||||||
|
} else {
|
||||||
|
t.Logf("verified: v4 receiver saw outer TOS 0x%02x (ECN=0x%02x) from dual-stack sender", gotTOS, gotTOS&0x03)
|
||||||
|
}
|
||||||
|
}
|
||||||
+12
-2
@@ -140,7 +140,7 @@ func (u *RIOConn) bind(l *slog.Logger, sa windows.Sockaddr) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *RIOConn) ListenOut(r EncReader) error {
|
func (u *RIOConn) ListenOut(r EncReader, flush func()) error {
|
||||||
buffer := make([]byte, MTU)
|
buffer := make([]byte, MTU)
|
||||||
|
|
||||||
var lastRecvErr time.Time
|
var lastRecvErr time.Time
|
||||||
@@ -161,7 +161,8 @@ func (u *RIOConn) ListenOut(r EncReader) error {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
r(netip.AddrPortFrom(netip.AddrFrom16(rua.Addr).Unmap(), (rua.Port>>8)|((rua.Port&0xff)<<8)), buffer[:n])
|
r(netip.AddrPortFrom(netip.AddrFrom16(rua.Addr).Unmap(), (rua.Port>>8)|((rua.Port&0xff)<<8)), buffer[:n], RxMeta{})
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -316,6 +317,15 @@ func (u *RIOConn) WriteTo(buf []byte, ip netip.AddrPort) error {
|
|||||||
return winrio.SendEx(u.rq, dataBuffer, 1, nil, addressBuffer, nil, nil, 0, 0)
|
return winrio.SendEx(u.rq, dataBuffer, 1, nil, addressBuffer, nil, nil, 0, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (u *RIOConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||||
|
for i, b := range bufs {
|
||||||
|
if err := u.WriteTo(b, addrs[i]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (u *RIOConn) LocalAddr() (netip.AddrPort, error) {
|
func (u *RIOConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
sa, err := windows.Getsockname(u.sock)
|
sa, err := windows.Getsockname(u.sock)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
+11
-2
@@ -157,15 +157,24 @@ func (u *TesterConn) WriteTo(b []byte, addr netip.AddrPort) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
func (u *TesterConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||||
|
for i, b := range bufs {
|
||||||
|
if err := u.WriteTo(b, addrs[i]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (u *TesterConn) ListenOut(r EncReader) error {
|
func (u *TesterConn) ListenOut(r EncReader, flush func()) error {
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
case <-u.done:
|
case <-u.done:
|
||||||
return os.ErrClosed
|
return os.ErrClosed
|
||||||
case p := <-u.RxPackets:
|
case p := <-u.RxPackets:
|
||||||
r(p.From, p.Data)
|
r(p.From, p.Data, RxMeta{})
|
||||||
p.Release()
|
p.Release()
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ import (
|
|||||||
// PinThreadToCPU restricts the calling OS thread to the given CPU via
|
// PinThreadToCPU restricts the calling OS thread to the given CPU via
|
||||||
// sched_setaffinity(2). Combined with runtime.LockOSThread on the
|
// sched_setaffinity(2). Combined with runtime.LockOSThread on the
|
||||||
// goroutine, this prevents the kernel from migrating us across CPUs and
|
// goroutine, this prevents the kernel from migrating us across CPUs and
|
||||||
// in turn keeps every UDP send from this goroutine going through the
|
// in turn keeps every sendmmsg from this goroutine going through the
|
||||||
// same XPS-selected TX ring, eliminating the wire-side reorder that
|
// same XPS-selected TX ring, eliminating the wire-side reorder that
|
||||||
// otherwise fragments one nebula flow across multiple rings.
|
// otherwise fragments one nebula flow across multiple rings.
|
||||||
func PinThreadToCPU(cpu int) error {
|
func PinThreadToCPU(cpu int) error {
|
||||||
|
|||||||
@@ -0,0 +1,198 @@
|
|||||||
|
//go:build linux && !android && !e2e_testing
|
||||||
|
|
||||||
|
package util
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// NICIRQCPUs returns the set of CPUs that service interrupts for the ACTIVE
|
||||||
|
// RX/TX queues of physical network interfaces that are up. Read-only: it
|
||||||
|
// matches /proc/interrupts action names against each NIC's PCI address and
|
||||||
|
// interface name, drops vectors whose queue index is beyond the device's
|
||||||
|
// active queue count (drivers like mlx5 keep handlers registered for
|
||||||
|
// deactivated queues, so /proc/interrupts alone over-reports), and unions
|
||||||
|
// /proc/irq/<n>/effective_affinity_list for the survivors.
|
||||||
|
//
|
||||||
|
// Callers use this to keep busy pinned threads OFF those CPUs: a thread
|
||||||
|
// pinned onto a core that also runs NAPI for a NIC RX queue competes with
|
||||||
|
// softirq processing for the core and measurably collapses throughput for
|
||||||
|
// flows hashed to that queue.
|
||||||
|
func NICIRQCPUs() (map[int]bool, error) {
|
||||||
|
return nicIRQCPUs("/sys/class/net", "/proc/irq", "/proc/interrupts")
|
||||||
|
}
|
||||||
|
|
||||||
|
// irqAction is one row of /proc/interrupts: the IRQ number and the action
|
||||||
|
// (handler) name in its final column, e.g. "mlx5_comp3@pci:0000:82:00.0".
|
||||||
|
type irqAction struct {
|
||||||
|
irq string
|
||||||
|
action string
|
||||||
|
}
|
||||||
|
|
||||||
|
func nicIRQCPUs(netDir, irqDir, interruptsPath string) (map[int]bool, error) {
|
||||||
|
actions, err := parseInterrupts(interruptsPath)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
devs, err := os.ReadDir(netDir)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
cpus := make(map[int]bool)
|
||||||
|
for _, dev := range devs {
|
||||||
|
devPath := filepath.Join(netDir, dev.Name())
|
||||||
|
pciDev, err := filepath.EvalSymlinks(filepath.Join(devPath, "device"))
|
||||||
|
if err != nil {
|
||||||
|
continue // virtual device (lo, tun, bridge, vlan, ...)
|
||||||
|
}
|
||||||
|
pciAddr := filepath.Base(pciDev)
|
||||||
|
state, err := os.ReadFile(filepath.Join(devPath, "operstate"))
|
||||||
|
if err != nil || strings.TrimSpace(string(state)) != "up" {
|
||||||
|
continue // a down NIC's queue IRQs don't fire
|
||||||
|
}
|
||||||
|
nq := countQueues(filepath.Join(devPath, "queues"))
|
||||||
|
|
||||||
|
for _, ia := range actions {
|
||||||
|
if !strings.Contains(ia.action, pciAddr) && !containsWord(ia.action, dev.Name()) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// Vector naming puts the queue index at the end of the handler
|
||||||
|
// name (mlx5_comp3@pci:..., ice-eth0-TxRx-3, virtio0-input.3).
|
||||||
|
// An index at or beyond the active queue count is a handler for
|
||||||
|
// a deactivated queue: registered, but it will not fire.
|
||||||
|
name, _, _ := strings.Cut(ia.action, "@")
|
||||||
|
if idx, ok := trailingInt(name); ok && idx >= nq {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for _, cpu := range irqAffinity(irqDir, ia.irq) {
|
||||||
|
cpus[cpu] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return cpus, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseInterrupts extracts (irq, action) pairs from /proc/interrupts,
|
||||||
|
// skipping the header and the non-numeric summary rows (NMI, LOC, ...).
|
||||||
|
func parseInterrupts(path string) ([]irqAction, error) {
|
||||||
|
data, err := os.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
var out []irqAction
|
||||||
|
for line := range strings.SplitSeq(string(data), "\n") {
|
||||||
|
fields := strings.Fields(line)
|
||||||
|
if len(fields) < 2 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
irq, ok := strings.CutSuffix(fields[0], ":")
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if _, err := strconv.Atoi(irq); err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
out = append(out, irqAction{irq: irq, action: fields[len(fields)-1]})
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// irqAffinity returns the CPUs IRQ n actually targets.
|
||||||
|
// effective_affinity_list is the vector's real target; smp_affinity_list
|
||||||
|
// (the fallback for kernels without effective affinity reporting) is the
|
||||||
|
// admin-allowed mask and may be wider.
|
||||||
|
func irqAffinity(irqDir, irq string) []int {
|
||||||
|
irqPath := filepath.Join(irqDir, irq)
|
||||||
|
list, err := os.ReadFile(filepath.Join(irqPath, "effective_affinity_list"))
|
||||||
|
if err != nil || len(strings.TrimSpace(string(list))) == 0 {
|
||||||
|
list, err = os.ReadFile(filepath.Join(irqPath, "smp_affinity_list"))
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return parseCPUList(strings.TrimSpace(string(list)))
|
||||||
|
}
|
||||||
|
|
||||||
|
// countQueues counts the rx-* entries of a netdev's queues directory — the
|
||||||
|
// device's ACTIVE RX queues (sysfs removes the directories when a queue is
|
||||||
|
// deactivated, e.g. by ethtool -L).
|
||||||
|
func countQueues(queuesDir string) int {
|
||||||
|
entries, err := os.ReadDir(queuesDir)
|
||||||
|
if err != nil {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
n := 0
|
||||||
|
for _, e := range entries {
|
||||||
|
if strings.HasPrefix(e.Name(), "rx-") {
|
||||||
|
n++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
// containsWord reports whether s contains word bounded by non-alphanumeric
|
||||||
|
// characters (or string edges), so ifname "eth0" doesn't match "eth01".
|
||||||
|
func containsWord(s, word string) bool {
|
||||||
|
for start := 0; ; {
|
||||||
|
i := strings.Index(s[start:], word)
|
||||||
|
if i < 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
i += start
|
||||||
|
before := i == 0 || !isAlnum(s[i-1])
|
||||||
|
afterIdx := i + len(word)
|
||||||
|
after := afterIdx == len(s) || !isAlnum(s[afterIdx])
|
||||||
|
if before && after {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
start = i + 1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func isAlnum(b byte) bool {
|
||||||
|
return b >= '0' && b <= '9' || b >= 'a' && b <= 'z' || b >= 'A' && b <= 'Z'
|
||||||
|
}
|
||||||
|
|
||||||
|
// trailingInt parses the decimal digits at the end of s.
|
||||||
|
func trailingInt(s string) (int, bool) {
|
||||||
|
i := len(s)
|
||||||
|
for i > 0 && s[i-1] >= '0' && s[i-1] <= '9' {
|
||||||
|
i--
|
||||||
|
}
|
||||||
|
if i == len(s) {
|
||||||
|
return 0, false
|
||||||
|
}
|
||||||
|
n, err := strconv.Atoi(s[i:])
|
||||||
|
return n, err == nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseCPUList parses the kernel's cpulist format: comma-separated CPU ids
|
||||||
|
// or inclusive ranges, e.g. "0-3,8,10-12". Malformed elements are skipped —
|
||||||
|
// this parses trusted kernel output, not user input.
|
||||||
|
func parseCPUList(s string) []int {
|
||||||
|
if s == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
var cpus []int
|
||||||
|
for part := range strings.SplitSeq(s, ",") {
|
||||||
|
lo, hi, ok := strings.Cut(part, "-")
|
||||||
|
start, err := strconv.Atoi(strings.TrimSpace(lo))
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
end := start
|
||||||
|
if ok {
|
||||||
|
if end, err = strconv.Atoi(strings.TrimSpace(hi)); err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for cpu := start; cpu <= end; cpu++ {
|
||||||
|
cpus = append(cpus, cpu)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return cpus
|
||||||
|
}
|
||||||
@@ -0,0 +1,153 @@
|
|||||||
|
//go:build linux && !android && !e2e_testing
|
||||||
|
|
||||||
|
package util
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"reflect"
|
||||||
|
"strconv"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestParseCPUList(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
in string
|
||||||
|
want []int
|
||||||
|
}{
|
||||||
|
{"", nil},
|
||||||
|
{"3", []int{3}},
|
||||||
|
{"0-3", []int{0, 1, 2, 3}},
|
||||||
|
{"0-2,8,10-11", []int{0, 1, 2, 8, 10, 11}},
|
||||||
|
{"garbage", nil},
|
||||||
|
{"1,garbage,4", []int{1, 4}},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
if got := parseCPUList(c.in); !reflect.DeepEqual(got, c.want) {
|
||||||
|
t.Errorf("parseCPUList(%q) = %v, want %v", c.in, got, c.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTrailingInt(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
in string
|
||||||
|
want int
|
||||||
|
ok bool
|
||||||
|
}{
|
||||||
|
{"mlx5_comp12", 12, true},
|
||||||
|
{"ice-eth0-TxRx-3", 3, true},
|
||||||
|
{"virtio0-input.7", 7, true},
|
||||||
|
{"mlx5_async0", 0, true},
|
||||||
|
{"no-digits", 0, false},
|
||||||
|
{"", 0, false},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
got, ok := trailingInt(c.in)
|
||||||
|
if got != c.want || ok != c.ok {
|
||||||
|
t.Errorf("trailingInt(%q) = (%d, %v), want (%d, %v)", c.in, got, ok, c.want, c.ok)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestContainsWord(t *testing.T) {
|
||||||
|
if !containsWord("ice-eth0-TxRx-3", "eth0") {
|
||||||
|
t.Error("eth0 should match with boundaries")
|
||||||
|
}
|
||||||
|
if containsWord("ice-eth01-TxRx-3", "eth0") {
|
||||||
|
t.Error("eth0 must not match inside eth01")
|
||||||
|
}
|
||||||
|
if !containsWord("eth0", "eth0") {
|
||||||
|
t.Error("exact match should work")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// fakeNIC builds /sys/class/net/<name> with operstate, a device symlink to a
|
||||||
|
// PCI-address-named dir (physical NICs only), and nq rx queue directories.
|
||||||
|
func fakeNIC(t *testing.T, netDir, name, operstate, pciAddr string, nq int) {
|
||||||
|
t.Helper()
|
||||||
|
devPath := filepath.Join(netDir, name)
|
||||||
|
if err := os.MkdirAll(devPath, 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(filepath.Join(devPath, "operstate"), []byte(operstate+"\n"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if pciAddr == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
pciDir := filepath.Join(netDir, "..", "devices", pciAddr)
|
||||||
|
if err := os.MkdirAll(pciDir, 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := os.Symlink(pciDir, filepath.Join(devPath, "device")); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
for i := 0; i < nq; i++ {
|
||||||
|
if err := os.MkdirAll(filepath.Join(devPath, "queues", "rx-"+strconv.Itoa(i)), 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeIRQ(t *testing.T, irqDir, irq, affinity string) {
|
||||||
|
t.Helper()
|
||||||
|
p := filepath.Join(irqDir, irq)
|
||||||
|
if err := os.MkdirAll(p, 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(filepath.Join(p, "effective_affinity_list"), []byte(affinity+"\n"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNICIRQCPUs(t *testing.T) {
|
||||||
|
root := t.TempDir()
|
||||||
|
netDir := filepath.Join(root, "class", "net")
|
||||||
|
irqDir := filepath.Join(root, "irq")
|
||||||
|
if err := os.MkdirAll(netDir, 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// eth0: up, 2 active queues at 0000:82:00.0. comp0/comp1 active,
|
||||||
|
// comp2 is a deactivated queue's leftover handler, async0 always fires.
|
||||||
|
fakeNIC(t, netDir, "eth0", "up", "0000:82:00.0", 2)
|
||||||
|
// eth1: physical but down; its vectors must not count.
|
||||||
|
fakeNIC(t, netDir, "eth1", "down", "0000:83:00.0", 2)
|
||||||
|
// eth9: up, matched by ifname (intel-style action names), 1 queue.
|
||||||
|
fakeNIC(t, netDir, "eth9", "up", "0000:84:00.0", 1)
|
||||||
|
// nebula1: virtual, no device dir.
|
||||||
|
fakeNIC(t, netDir, "nebula1", "up", "", 0)
|
||||||
|
|
||||||
|
interrupts := filepath.Join(root, "interrupts")
|
||||||
|
content := ` CPU0 CPU1
|
||||||
|
100: 1 2 IR-PCI-MSIX 1-edge mlx5_comp0@pci:0000:82:00.0
|
||||||
|
101: 1 2 IR-PCI-MSIX 2-edge mlx5_comp1@pci:0000:82:00.0
|
||||||
|
102: 1 2 IR-PCI-MSIX 3-edge mlx5_comp2@pci:0000:82:00.0
|
||||||
|
103: 1 2 IR-PCI-MSIX 4-edge mlx5_async0@pci:0000:82:00.0
|
||||||
|
200: 1 2 IR-PCI-MSIX 5-edge mlx5_comp0@pci:0000:83:00.0
|
||||||
|
300: 1 2 IR-PCI-MSIX 6-edge ice-eth9-TxRx-0
|
||||||
|
301: 1 2 IR-PCI-MSIX 7-edge ice-eth9-TxRx-1
|
||||||
|
NMI: 0 0 Non-maskable interrupts
|
||||||
|
`
|
||||||
|
if err := os.WriteFile(interrupts, []byte(content), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
writeIRQ(t, irqDir, "100", "0-1") // eth0 comp0: counted
|
||||||
|
writeIRQ(t, irqDir, "101", "2") // eth0 comp1: counted
|
||||||
|
writeIRQ(t, irqDir, "102", "5") // eth0 comp2: beyond 2 queues, skipped
|
||||||
|
writeIRQ(t, irqDir, "103", "7") // eth0 async0: counted
|
||||||
|
writeIRQ(t, irqDir, "200", "9") // eth1 down: skipped
|
||||||
|
writeIRQ(t, irqDir, "300", "11") // eth9 TxRx-0: counted
|
||||||
|
writeIRQ(t, irqDir, "301", "12") // eth9 TxRx-1: beyond 1 queue, skipped
|
||||||
|
|
||||||
|
got, err := nicIRQCPUs(netDir, irqDir, interrupts)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("nicIRQCPUs: %v", err)
|
||||||
|
}
|
||||||
|
want := map[int]bool{0: true, 1: true, 2: true, 7: true, 11: true}
|
||||||
|
if !reflect.DeepEqual(got, want) {
|
||||||
|
t.Fatalf("nicIRQCPUs = %v, want %v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,9 @@
|
|||||||
|
//go:build !linux || android || e2e_testing
|
||||||
|
|
||||||
|
package util
|
||||||
|
|
||||||
|
// NICIRQCPUs reports no IRQ information on platforms without the linux
|
||||||
|
// sysfs interface; callers fall back to their non-IRQ-aware defaults.
|
||||||
|
func NICIRQCPUs() (map[int]bool, error) {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user