mirror of
https://github.com/slackhq/nebula.git
synced 2026-10-08 08:07:56 +02:00
the definitive tun offloads branch (#1704)
This commit is contained in:
@@ -0,0 +1,187 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"math/rand"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// The checksum-seeding helpers feed the virtio NEEDS_CSUM contract: the L4
|
||||
// checksum field is pre-loaded with the folded (not inverted) pseudo-header
|
||||
// sum, and the kernel later adds the L4 byte sum and inverts. A wrong seed
|
||||
// produces packets every receiver silently drops, with nothing failing on
|
||||
// our side — so these tests check the helpers against an independent
|
||||
// RFC 1071 reference built from explicit pseudo-header bytes, never against
|
||||
// the production checksum code.
|
||||
|
||||
// refSum accumulates big-endian 16-bit words of b (odd tail zero-padded)
|
||||
// into a wide one's-complement accumulator.
|
||||
func refSum(b []byte) uint64 {
|
||||
var s uint64
|
||||
for i := 0; i+1 < len(b); i += 2 {
|
||||
s += uint64(b[i])<<8 | uint64(b[i+1])
|
||||
}
|
||||
if len(b)%2 == 1 {
|
||||
s += uint64(b[len(b)-1]) << 8
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// refFold folds a wide one's-complement accumulator to 16 bits.
|
||||
func refFold(s uint64) uint16 {
|
||||
for s>>16 != 0 {
|
||||
s = s&0xffff + s>>16
|
||||
}
|
||||
return uint16(s)
|
||||
}
|
||||
|
||||
func TestFoldOnceNoInvertEdgeCases(t *testing.T) {
|
||||
cases := []uint32{
|
||||
0, 1, 0xffff,
|
||||
0x10000, // single carry
|
||||
0x1fffe, // 0xffff + 0xffff: carry produces another 0xffff
|
||||
0xffff0000, // high half only
|
||||
0xfffeffff, // fold yields 0x1fffd: needs a second fold
|
||||
0xffffffff, // worst case
|
||||
0x00010001, // simple two-word
|
||||
}
|
||||
for _, c := range cases {
|
||||
want := refFold(uint64(c))
|
||||
if got := foldOnceNoInvert(c); got != want {
|
||||
t.Errorf("foldOnceNoInvert(%#x) = %#x, want %#x", c, got, want)
|
||||
}
|
||||
// Folding a folded value must be a no-op.
|
||||
if got := foldOnceNoInvert(uint32(foldOnceNoInvert(c))); got != foldOnceNoInvert(c) {
|
||||
t.Errorf("foldOnceNoInvert not idempotent at %#x", c)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPseudoSumIPv4MatchesReference(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
src, dst [4]byte
|
||||
proto byte
|
||||
l4Len int
|
||||
}{
|
||||
{"simple", [4]byte{10, 0, 0, 1}, [4]byte{10, 0, 0, 2}, 6, 20},
|
||||
{"zero-len", [4]byte{192, 168, 1, 1}, [4]byte{192, 168, 1, 2}, 17, 0},
|
||||
{"max-len", [4]byte{1, 2, 3, 4}, [4]byte{5, 6, 7, 8}, 6, 65535},
|
||||
{"carry-heavy", [4]byte{255, 255, 255, 255}, [4]byte{255, 255, 255, 254}, 17, 65535},
|
||||
{"broadcastish", [4]byte{255, 255, 255, 255}, [4]byte{255, 255, 255, 255}, 255, 65535},
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
// RFC 793 pseudo-header: src(4) dst(4) zero(1) proto(1) len(2).
|
||||
ph := make([]byte, 12)
|
||||
copy(ph[0:4], c.src[:])
|
||||
copy(ph[4:8], c.dst[:])
|
||||
ph[9] = c.proto
|
||||
binary.BigEndian.PutUint16(ph[10:12], uint16(c.l4Len))
|
||||
want := refFold(refSum(ph))
|
||||
|
||||
got := foldOnceNoInvert(pseudoSumIPv4(c.src[:], c.dst[:], c.proto, c.l4Len))
|
||||
if got != want {
|
||||
t.Errorf("fold(pseudoSumIPv4) = %#x, want %#x", got, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPseudoSumIPv6MatchesReference(t *testing.T) {
|
||||
ones := func(b byte) (a [16]byte) {
|
||||
for i := range a {
|
||||
a[i] = b
|
||||
}
|
||||
return
|
||||
}
|
||||
cases := []struct {
|
||||
name string
|
||||
src, dst [16]byte
|
||||
proto byte
|
||||
l4Len int
|
||||
}{
|
||||
{"simple", [16]byte{0xfe, 0x80, 15: 1}, [16]byte{0xfe, 0x80, 15: 2}, 6, 20},
|
||||
{"zero-len", [16]byte{0x20, 0x01, 15: 9}, [16]byte{0x20, 0x01, 15: 8}, 17, 0},
|
||||
{"max-u16-len", ones(0xff), ones(0xfe), 6, 65535},
|
||||
{"len-past-u16", ones(0xff), ones(0xff), 17, 0x12345}, // exercises the 32-bit split
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
// RFC 8200 pseudo-header: src(16) dst(16) len(4) zero(3) next(1).
|
||||
ph := make([]byte, 40)
|
||||
copy(ph[0:16], c.src[:])
|
||||
copy(ph[16:32], c.dst[:])
|
||||
binary.BigEndian.PutUint32(ph[32:36], uint32(c.l4Len))
|
||||
ph[39] = c.proto
|
||||
want := refFold(refSum(ph))
|
||||
|
||||
got := foldOnceNoInvert(pseudoSumIPv6(c.src[:], c.dst[:], c.proto, c.l4Len))
|
||||
if got != want {
|
||||
t.Errorf("fold(pseudoSumIPv6) = %#x, want %#x", got, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIPv4HdrChecksumMatchesReference(t *testing.T) {
|
||||
rng := rand.New(rand.NewSource(0x1791))
|
||||
for _, hdrLen := range []int{20, 24, 40, 60} {
|
||||
for trial := 0; trial < 200; trial++ {
|
||||
hdr := make([]byte, hdrLen)
|
||||
rng.Read(hdr)
|
||||
hdr[0] = 0x40 | byte(hdrLen/4)
|
||||
hdr[10], hdr[11] = 0, 0 // checksum field zeroed, as the contract requires
|
||||
|
||||
want := ^refFold(refSum(hdr))
|
||||
got := ipv4HdrChecksum(hdr)
|
||||
if got != want {
|
||||
t.Fatalf("ipv4HdrChecksum(len=%d trial=%d) = %#x, want %#x", hdrLen, trial, got, want)
|
||||
}
|
||||
|
||||
// Receiver-side property: with the checksum stored, the full
|
||||
// header must sum to all-ones.
|
||||
binary.BigEndian.PutUint16(hdr[10:12], got)
|
||||
if v := refFold(refSum(hdr)); v != 0xffff {
|
||||
t.Fatalf("stored checksum does not validate: full-header fold = %#x", v)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestChecksumSeedReceiverAcceptance is the end-to-end property the helpers
|
||||
// exist for: seed the TCP checksum field with fold(pseudoSum), do what the
|
||||
// kernel's NEEDS_CSUM completion does (one's-complement sum over the L4
|
||||
// bytes including the seed, then invert, then store), and verify the result
|
||||
// the way a receiver does (pseudo-header + L4 must sum to all-ones).
|
||||
func TestChecksumSeedReceiverAcceptance(t *testing.T) {
|
||||
rng := rand.New(rand.NewSource(0x1826))
|
||||
for trial := 0; trial < 200; trial++ {
|
||||
src := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
|
||||
dst := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
|
||||
payLen := rng.Intn(1500)
|
||||
l4 := make([]byte, 20+payLen)
|
||||
rng.Read(l4)
|
||||
|
||||
// Seed exactly as flushSlot does.
|
||||
seed := foldOnceNoInvert(pseudoSumIPv4(src[:], dst[:], 6, len(l4)))
|
||||
binary.BigEndian.PutUint16(l4[16:18], seed)
|
||||
|
||||
// Kernel NEEDS_CSUM completion: sum the L4 region (seed included,
|
||||
// which is equivalent to summing with the field zeroed and folding
|
||||
// the seed in), invert, store.
|
||||
final := ^refFold(refSum(l4[:16]) + uint64(seed) + refSum(l4[18:]))
|
||||
binary.BigEndian.PutUint16(l4[16:18], final)
|
||||
|
||||
// Receiver validation.
|
||||
ph := make([]byte, 12)
|
||||
copy(ph[0:4], src[:])
|
||||
copy(ph[4:8], dst[:])
|
||||
ph[9] = 6
|
||||
binary.BigEndian.PutUint16(ph[10:12], uint16(len(l4)))
|
||||
if v := refFold(refSum(ph) + refSum(l4)); v != 0xffff {
|
||||
t.Fatalf("trial %d: receiver rejects packet: fold = %#x (seed=%#x final=%#x payLen=%d)",
|
||||
trial, v, seed, final, payLen)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,169 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
)
|
||||
|
||||
// SortKey identifies a packet's position in its sender's transmission order.
|
||||
type SortKey struct {
|
||||
// Epoch is a receiver-local ordinal for the tunnel (ConnectionState) that decrypted the packet:
|
||||
// a re-handshake replaces the tunnel outright and the replacement's epoch is higher,
|
||||
// so the old tunnel's packets sort first during the cutover overlap.
|
||||
Epoch uint64
|
||||
// Counter is the packet's AEAD message counter within that tunnel.
|
||||
Counter uint64
|
||||
}
|
||||
|
||||
// flowKey identifies a transport flow by {src, dst, sport, dport, family}.
|
||||
// Comparable, so map lookups and linear scans over the slot list stay tight.
|
||||
// Shared by the TCP and UDP coalescers; each coalescer keeps its own
|
||||
// openSlots map, so a TCP and UDP flow on the same 5-tuple-without-proto never alias.
|
||||
type flowKey struct {
|
||||
src, dst [16]byte
|
||||
sport, dport uint16
|
||||
isV6 bool
|
||||
}
|
||||
|
||||
// initialSlots is the starting capacity of the slot pool.
|
||||
// One flow per packet is the worst case, so this matches a typical carrier-side recvmmsg batch on the UDP socket.
|
||||
const initialSlots = 64
|
||||
|
||||
// parseIPAt validates the IP header for lane parsing. newPacket already resolved the L4 protocol
|
||||
// and offset for the firewall, so there is no proto sniff here; the caller's ipHdrLen is
|
||||
// cross-checked instead. A plain header (v4 IHL 20, v6 exactly 40) is the only coalesceable
|
||||
// shape. The v6 check is load-bearing: it rejects extension-header packets whose L4 is not at byte 40.
|
||||
//
|
||||
// The prologues fill fk's addresses and family in place (ports belong to the L4 parser; fk must
|
||||
// be zero on entry so the v4 path leaves src[4:]/dst[4:] clear for map equality) and return pkt
|
||||
// trimmed to the IP-declared length. The receiver-as-out-pointer shape is deliberate: these
|
||||
// functions are too big to inline, and returning structs by value put five 64-byte copies on the
|
||||
// per-packet path.
|
||||
func (fk *flowKey) parseIPAt(pkt []byte, ipHdrLen int) ([]byte, bool) {
|
||||
if len(pkt) < 20 {
|
||||
return nil, false
|
||||
}
|
||||
switch pkt[0] >> 4 {
|
||||
case 4:
|
||||
if ipHdrLen != 20 {
|
||||
return nil, false
|
||||
}
|
||||
return fk.parseIPv4Prologue(pkt)
|
||||
case 6:
|
||||
if ipHdrLen != 40 || len(pkt) < 40 {
|
||||
return nil, false
|
||||
}
|
||||
return fk.parseIPv6Prologue(pkt)
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// parseIPv4Prologue is the shared IPv4 tail of the prologue entries; the caller has verified
|
||||
// len(pkt) >= 20 and the version.
|
||||
func (fk *flowKey) parseIPv4Prologue(pkt []byte) ([]byte, bool) {
|
||||
ihl := int(pkt[0]&0x0f) * 4
|
||||
if ihl != 20 {
|
||||
return nil, false
|
||||
}
|
||||
// Reject any fragmentation (MF or nonzero offset). The dispatcher already gated FragAny; kept
|
||||
// as defense in depth, since a fragment folded into a superpacket would corrupt reassembly.
|
||||
if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 {
|
||||
return nil, false
|
||||
}
|
||||
totalLen := int(binary.BigEndian.Uint16(pkt[2:4]))
|
||||
if totalLen > len(pkt) || totalLen < ihl {
|
||||
return nil, false
|
||||
}
|
||||
fk.isV6 = false
|
||||
copy(fk.src[:4], pkt[12:16])
|
||||
copy(fk.dst[:4], pkt[16:20])
|
||||
return pkt[:totalLen], true
|
||||
}
|
||||
|
||||
// parseIPv6Prologue is the shared IPv6 tail; the caller has verified len(pkt) >= 40, the version,
|
||||
// and that the L4 header sits at byte 40.
|
||||
func (fk *flowKey) parseIPv6Prologue(pkt []byte) ([]byte, bool) {
|
||||
payloadLen := int(binary.BigEndian.Uint16(pkt[4:6]))
|
||||
if 40+payloadLen > len(pkt) {
|
||||
return nil, false
|
||||
}
|
||||
fk.isV6 = true
|
||||
copy(fk.src[:], pkt[8:24])
|
||||
copy(fk.dst[:], pkt[24:40])
|
||||
return pkt[:40+payloadLen], true
|
||||
}
|
||||
|
||||
// ipHeadersMatch compares the IP portion of two packet header prefixes for
|
||||
// byte-for-byte equality on every field that must be identical across coalesced segments.
|
||||
// Size/IPID/IPCsum are masked out.
|
||||
// The full DSCP/ECN byte (IPv4 ToS / IPv6 traffic class) is compared, matching Linux kernel GRO:
|
||||
// segments with differing ECN codepoints must not coalesce,
|
||||
// otherwise ORing e.g. ECT(0) with ECT(1) would fabricate a false CE (congestion) mark or mark a Not-ECT flow as ECN-capable.
|
||||
//
|
||||
// The transport (L4) portion of the header is checked separately by the per-protocol matcher.
|
||||
func ipHeadersMatch(a, b []byte, isV6 bool) bool {
|
||||
if isV6 {
|
||||
// IPv6: [0:4] = version/TC/flow label (TC[1:0] is ECN, so the full TC byte must match),
|
||||
// [6:40] = next_hdr/hop + src + dst. Skip [4:6] payload_len.
|
||||
return bytes.Equal(a[:4], b[:4]) && bytes.Equal(a[6:40], b[6:40])
|
||||
}
|
||||
// IPv4: [0:2] = version/IHL + DSCP|ECN (full ECN byte must match),
|
||||
// [6:10] = flags/fragoff/TTL/proto, [12:20] = src+dst.
|
||||
// Skip [2:4] total len, [4:6] id, [10:12] csum.
|
||||
return bytes.Equal(a[:2], b[:2]) && bytes.Equal(a[6:10], b[6:10]) && bytes.Equal(a[12:20], b[12:20])
|
||||
}
|
||||
|
||||
// ipv4FlagDF is the Don't Fragment bit in the IPv4 flags byte (header byte 6).
|
||||
const ipv4FlagDF = 0x40
|
||||
|
||||
// ipv4CanCoalesceID reports whether an IPv4 packet whose header starts at
|
||||
// nextHdr may join a chain whose seed header is seedHdr as segment index seg
|
||||
// (the seed is segment 0). Kernel GSO re-stamps outgoing segment IDs as
|
||||
// seed_id+n, so coalescing is only transparent when that re-stamp is either
|
||||
// harmless (DF set: RFC 6864 atomic datagrams, the ID carries no meaning) or
|
||||
// reproduces the original IDs exactly (DF clear + IDs already sequential —
|
||||
// the same admission rule kernel GRO applies). Without this, a DF=0 sender
|
||||
// with non-sequential IDs (e.g. OpenBSD's randomized IDs) could have IDs
|
||||
// rewritten into ranges that collide across superpackets, corrupting
|
||||
// reassembly if the packets are fragmented after the TUN write.
|
||||
//
|
||||
// DF itself is guaranteed uniform across a chain by ipHeadersMatch (byte 6
|
||||
// is inside its compared range), so checking the seed's copy suffices.
|
||||
func ipv4CanCoalesceID(seedHdr, nextHdr []byte, seg int) bool {
|
||||
if seedHdr[6]&ipv4FlagDF != 0 {
|
||||
return true
|
||||
}
|
||||
expect := binary.BigEndian.Uint16(seedHdr[4:6]) + uint16(seg)
|
||||
return binary.BigEndian.Uint16(nextHdr[4:6]) == expect
|
||||
}
|
||||
|
||||
// Arena is an injectable byte-slab that hands out non-overlapping borrowed
|
||||
// slices via Reserve and releases them in bulk via Reset.
|
||||
type Arena struct {
|
||||
buf []byte
|
||||
}
|
||||
|
||||
// NewArena returns an Arena with a pre-allocated backing of the given capacity.
|
||||
func NewArena(capacity int) *Arena {
|
||||
return &Arena{buf: make([]byte, 0, capacity)}
|
||||
}
|
||||
|
||||
// Reserve hands out a non-overlapping sz-byte slice from the arena.
|
||||
// If the request doesn't fit the current backing, a fresh, larger backing is allocated.
|
||||
// Already-borrowed slices reference the old backing and remain valid until Reset.
|
||||
func (a *Arena) Reserve(sz int) []byte {
|
||||
if len(a.buf)+sz > cap(a.buf) {
|
||||
newCap := max(cap(a.buf)*2, sz)
|
||||
a.buf = make([]byte, 0, newCap)
|
||||
}
|
||||
start := len(a.buf)
|
||||
a.buf = a.buf[:start+sz]
|
||||
return a.buf[start : start+sz : start+sz]
|
||||
}
|
||||
|
||||
// Reset releases every slice handed out since the last Reset.
|
||||
// Callers must not use any previously-borrowed slice after this returns.
|
||||
// The underlying backing array is retained so subsequent Reserves don't re-allocate.
|
||||
func (a *Arena) Reset() {
|
||||
a.buf = a.buf[:0]
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/test"
|
||||
)
|
||||
|
||||
// stagePackets builds the stagedPacket entries Commit would have produced, so dispatch benchmarks
|
||||
// bypass staging and the sort entirely.
|
||||
func stagePackets(pkts [][]byte) []stagedPacket {
|
||||
staged := make([]stagedPacket, len(pkts))
|
||||
for i, p := range pkts {
|
||||
pp := testPP(p)
|
||||
staged[i] = stagedPacket{
|
||||
pkt: p,
|
||||
key: SortKey{Epoch: 1, Counter: uint64(i + 1)},
|
||||
proto: pp.Protocol,
|
||||
fragAny: pp.FragAny,
|
||||
ipHdrLen: uint16(pp.IPHdrLen),
|
||||
}
|
||||
}
|
||||
return staged
|
||||
}
|
||||
|
||||
func flushLanes(b *testing.B, m *MultiCoalescer) {
|
||||
b.Helper()
|
||||
if m.tcp != nil {
|
||||
if err := m.tcp.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
if m.udp != nil {
|
||||
if err := m.udp.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := m.pt.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
// runDispatchBench measures dispatch plus the per-batch lane flush: the post-sort half of the
|
||||
// batcher, which is where the production profile concentrates.
|
||||
func runDispatchBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||
b.Helper()
|
||||
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger())
|
||||
staged := stagePackets(pkts)
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkts[0])))
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
if err := m.dispatch(staged[i%len(staged)]); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if (i+1)%batchSize == 0 {
|
||||
flushLanes(b, m)
|
||||
}
|
||||
}
|
||||
b.StopTimer()
|
||||
flushLanes(b, m)
|
||||
}
|
||||
|
||||
// BenchmarkDispatchSingleFlow is the bulk steady state: every packet past the seed appends.
|
||||
func BenchmarkDispatchSingleFlow(b *testing.B) {
|
||||
runDispatchBench(b, buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200), tcpCoalesceMaxSegs)
|
||||
}
|
||||
|
||||
// BenchmarkDispatchInterleaved16 stresses the openSlots map: 16 flows round-robined defeats the
|
||||
// lastSlot cache on every packet.
|
||||
func BenchmarkDispatchInterleaved16(b *testing.B) {
|
||||
pkts := buildTCPv4Interleaved(16, tcpCoalesceMaxSegs, 1200)
|
||||
runDispatchBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkDispatchAckHeavy alternates MSS data with pure ACKs on one flow — the RX shape of a
|
||||
// bidirectional transfer (the peer's data and its ACKs of our data share the tunnel direction).
|
||||
func BenchmarkDispatchAckHeavy(b *testing.B) {
|
||||
pay := make([]byte, 1200)
|
||||
var pkts [][]byte
|
||||
seq := uint32(1000)
|
||||
for range tcpCoalesceMaxSegs / 2 {
|
||||
pkts = append(pkts, buildTCPv4(seq, tcpAck, pay))
|
||||
seq += uint32(len(pay))
|
||||
pkts = append(pkts, buildTCPv4(seq, tcpAck, nil))
|
||||
}
|
||||
runDispatchBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkDispatchUDPFlow is the QUIC-ish bulk UDP shape.
|
||||
func BenchmarkDispatchUDPFlow(b *testing.B) {
|
||||
pay := make([]byte, 1200)
|
||||
pkts := make([][]byte, udpCoalesceMaxSegs)
|
||||
for i := range pkts {
|
||||
pkts[i] = buildUDPv4(2000, 443, pay)
|
||||
}
|
||||
runDispatchBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkDispatchSeedHeavy sets PSH on every packet so each one seeds and immediately closes
|
||||
// its own slot — the small-write RPC shape, and the upper bound on what the seed path (including
|
||||
// the parsedTCP-to-slot field transfer) can cost.
|
||||
func BenchmarkDispatchSeedHeavy(b *testing.B) {
|
||||
pay := make([]byte, 1200)
|
||||
pkts := make([][]byte, tcpCoalesceMaxSegs)
|
||||
seq := uint32(1000)
|
||||
for i := range pkts {
|
||||
pkts[i] = buildTCPv4(seq, tcpAckPsh, pay)
|
||||
seq += uint32(len(pay))
|
||||
}
|
||||
runDispatchBench(b, pkts, len(pkts))
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
package batch
|
||||
|
||||
//TODO refactor this away
|
||||
// This file holds the lanes' self-parsing Commit entries and the proto-checking parsers behind
|
||||
// them. Production traffic enters the lanes only through MultiCoalescer.dispatch and the At
|
||||
// parsers; these wrappers reproduce that path (including seal-all on unparseable shapes) on top
|
||||
// of a local parse, so tests and benches can drive one lane with nothing but a packet.
|
||||
|
||||
// parseIPPrologue resolves the IP version, requires the L4 protocol to match wantProto (6 TCP,
|
||||
// 17 UDP), and defers to the shared per-version cores. Returns the trimmed packet and the L4
|
||||
// offset; fk must be zero on entry and is filled in place.
|
||||
func (fk *flowKey) parseIPPrologue(pkt []byte, wantProto byte) ([]byte, int, bool) {
|
||||
if len(pkt) < 20 {
|
||||
return nil, 0, false
|
||||
}
|
||||
switch pkt[0] >> 4 {
|
||||
case 4:
|
||||
if pkt[9] != wantProto {
|
||||
return nil, 0, false
|
||||
}
|
||||
trimmed, ok := fk.parseIPv4Prologue(pkt)
|
||||
return trimmed, 20, ok
|
||||
case 6:
|
||||
if len(pkt) < 40 {
|
||||
return nil, 0, false
|
||||
}
|
||||
if pkt[6] != wantProto {
|
||||
return nil, 0, false
|
||||
}
|
||||
trimmed, ok := fk.parseIPv6Prologue(pkt)
|
||||
return trimmed, 40, ok
|
||||
}
|
||||
return nil, 0, false
|
||||
}
|
||||
|
||||
// parseBase extracts the flow key and IP/TCP offsets for any TCP packet, admissible for
|
||||
// coalescing or not. Returns false for non-TCP or malformed input.
|
||||
func (p *parsedTCP) parseBase(pkt []byte) bool {
|
||||
trimmed, ipHdrLen, ok := p.fk.parseIPPrologue(pkt, ipProtoTCP)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return p.parseTail(trimmed, ipHdrLen)
|
||||
}
|
||||
|
||||
// parseBase extracts the flow key and IP/UDP offsets for a UDP packet.
|
||||
func (p *parsedUDP) parseBase(pkt []byte) bool {
|
||||
trimmed, ipHdrLen, ok := p.fk.parseIPPrologue(pkt, ipProtoUDP)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return p.parseTail(trimmed, ipHdrLen)
|
||||
}
|
||||
|
||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
||||
func (c *TCPCoalescer) Commit(pkt []byte) error {
|
||||
var info parsedTCP
|
||||
if !info.parseBase(pkt) {
|
||||
// Unparseable: flow key unknown, seal everything so later data cannot emit ahead of it.
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(pkt)
|
||||
return nil
|
||||
}
|
||||
return c.commitParsed(pkt, &info)
|
||||
}
|
||||
|
||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
||||
func (c *UDPCoalescer) Commit(pkt []byte) error {
|
||||
var info parsedUDP
|
||||
if !info.parseBase(pkt) {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(pkt)
|
||||
return nil
|
||||
}
|
||||
return c.commitParsed(pkt, &info)
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"cmp"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"slices"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
)
|
||||
|
||||
// MultiCoalescer stages plaintext packets with their (epoch, counter) sort keys and, at Flush,
|
||||
// replays them in sender-transmission order into lane-specific batchers selected by L4 protocol.
|
||||
//
|
||||
// Sorting before dispatch keeps the ordering story simple: each lane consumes packets in
|
||||
// transmission order, builds slots in that order, and emits them in creation order. Wire reorder
|
||||
// inside a flush batch is repaired here, before it can fragment a lane's coalesce chains, so the
|
||||
// lanes carry no reorder-repair machinery.
|
||||
//
|
||||
// The contract is per-tunnel transmission order within each lane, with two exceptions: a pure TCP
|
||||
// ACK may be overtaken by later same-flow data (it does not close the flow's open chain; a late
|
||||
// ACK is just a stale ACK), and an unparseable shape seals every open chain in its lane (its flow
|
||||
// is unknown) and rides the lane as an in-lane verbatim, still in transmission order. Routing
|
||||
// follows the flow: a flow's non-coalesceable shapes ride its protocol lane rather than falling
|
||||
// to the later-flushed pt lane.
|
||||
//
|
||||
// Cross-lane order (TCP vs UDP vs everything else) is not preserved.
|
||||
type MultiCoalescer struct {
|
||||
tcp *TCPCoalescer
|
||||
udp *UDPCoalescer
|
||||
pt *Passthrough
|
||||
|
||||
// staged holds this batch's packets and sort keys until Flush. Borrowed: the caller keeps
|
||||
// each pkt alive until Flush returns.
|
||||
staged []stagedPacket
|
||||
}
|
||||
|
||||
// stagedPacket carries the scalars dispatch needs from the firewall's ParsedPacket, copied by
|
||||
// value: pp is reused by the caller per packet and must not be retained past Commit.
|
||||
type stagedPacket struct {
|
||||
pkt []byte
|
||||
key SortKey
|
||||
proto byte
|
||||
fragAny bool
|
||||
ipHdrLen uint16
|
||||
}
|
||||
|
||||
// NewMultiCoalescer builds a multi-lane batcher over w, based on available protocol support. The
|
||||
// staging sort applies even when no GSO lane is available: passthrough-only platforms still get
|
||||
// transmission-order repair.
|
||||
func NewMultiCoalescer(w io.Writer, l *slog.Logger) *MultiCoalescer {
|
||||
m := &MultiCoalescer{
|
||||
pt: NewPassthrough(w),
|
||||
staged: make([]stagedPacket, 0, initialSlots),
|
||||
}
|
||||
m.tcp = NewTCPCoalescer(w, l)
|
||||
m.udp = NewUDPCoalescer(w)
|
||||
return m
|
||||
}
|
||||
|
||||
// Commit stages pkt for the next Flush; dispatch is deferred so it runs on packets already in
|
||||
// transmission order. key carries the packet's tunnel epoch and message counter. pkt is borrowed:
|
||||
// the caller must keep it valid until the next Flush and not re-use it, and Flush may patch a
|
||||
// coalesced packet's headers in place. pp is the firewall's parse of pkt and is borrowed only
|
||||
// for this call, so the fields dispatch needs are copied here.
|
||||
func (m *MultiCoalescer) Commit(pkt []byte, key SortKey, pp *firewall.ParsedPacket) error {
|
||||
m.staged = append(m.staged, stagedPacket{
|
||||
pkt: pkt,
|
||||
key: key,
|
||||
proto: pp.Protocol,
|
||||
fragAny: pp.FragAny,
|
||||
ipHdrLen: uint16(pp.IPHdrLen),
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
// compareStaged orders staged packets by (epoch, counter)
|
||||
func compareStaged(a, b stagedPacket) int {
|
||||
if c := cmp.Compare(a.key.Epoch, b.key.Epoch); c != 0 {
|
||||
return c
|
||||
}
|
||||
return cmp.Compare(a.key.Counter, b.key.Counter)
|
||||
}
|
||||
|
||||
// dispatch routes one staged packet to its protocol lane (see commitStaged), or to the verbatim
|
||||
// passthrough when the lane has no GSO support.
|
||||
func (m *MultiCoalescer) dispatch(sp stagedPacket) error {
|
||||
switch sp.proto {
|
||||
case ipProtoTCP:
|
||||
if m.tcp != nil {
|
||||
return m.tcp.commitStaged(sp)
|
||||
}
|
||||
case ipProtoUDP:
|
||||
if m.udp != nil {
|
||||
return m.udp.commitStaged(sp)
|
||||
}
|
||||
}
|
||||
return m.pt.enqueue(sp.pkt)
|
||||
}
|
||||
|
||||
// Flush sorts the staged batch into transmission order, replays it into the lanes, then flushes each lane.
|
||||
// Drains everything and returns the joined errors; one bad packet does not hold up the rest.
|
||||
// After Flush returns, committed payload slices may be recycled.
|
||||
func (m *MultiCoalescer) Flush() error {
|
||||
// Arrival order is already almost sorted (reorder is the exception), which pdqsort detects
|
||||
// and handles in near-linear time.
|
||||
slices.SortFunc(m.staged, compareStaged)
|
||||
|
||||
var errs []error
|
||||
for _, sp := range m.staged {
|
||||
if err := m.dispatch(sp); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
clear(m.staged) // drop borrowed pkt refs
|
||||
m.staged = m.staged[:0]
|
||||
|
||||
if m.tcp != nil {
|
||||
if err := m.tcp.Flush(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
if m.udp != nil {
|
||||
if err := m.udp.Flush(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
if err := m.pt.Flush(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
@@ -0,0 +1,437 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/test"
|
||||
)
|
||||
|
||||
// keySeq hands out SortKeys with ascending counters in a fixed epoch, for
|
||||
// tests where commit order IS transmission order.
|
||||
type keySeq struct {
|
||||
epoch, counter uint64
|
||||
}
|
||||
|
||||
func (k *keySeq) next() SortKey {
|
||||
k.counter++
|
||||
return SortKey{Epoch: k.epoch, Counter: k.counter}
|
||||
}
|
||||
|
||||
// newTestMultiCoalescer builds a batcher over w.
|
||||
func newTestMultiCoalescer(tb testing.TB, w io.Writer) *MultiCoalescer {
|
||||
tb.Helper()
|
||||
return NewMultiCoalescer(w, test.NewLogger())
|
||||
}
|
||||
|
||||
// TestMultiCoalescerRoutesByProto confirms TCP/UDP/other land in the right
|
||||
// lane: TCP and UDP get coalesced when their lanes are enabled, anything
|
||||
// else (ICMP here) falls through to plain Write.
|
||||
func TestMultiCoalescerRoutesByProto(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
|
||||
tcpPay := make([]byte, 1200)
|
||||
udpPay := make([]byte, 1200)
|
||||
icmp := make([]byte, 28)
|
||||
icmp[0] = 0x45
|
||||
icmp[2] = 0
|
||||
icmp[3] = 28
|
||||
icmp[9] = 1
|
||||
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(1000, tcpAck, tcpPay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(2200, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(2200, tcpAck, tcpPay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(icmp, k.next(), testPP(icmp)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 1 TCP super (2 segments) + 1 UDP super (2 segments) = 2 gso writes.
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 gso writes (one TCP + one UDP), got %d", len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes) != 1 {
|
||||
t.Fatalf("want 1 plain write (ICMP), got %d", len(w.writes))
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerRestoresTransmissionOrder is the core staging-sort
|
||||
// property: packets committed out of counter order (wire reorder inside one
|
||||
// flush batch) are replayed into the lanes in transmission order, so the
|
||||
// reorder never fragments the coalesce chain — one superpacket, in seq
|
||||
// order, exactly as if the wire had never reordered. The retransmit shape
|
||||
// falls out of the same key: a retransmit carries a lower seq but a HIGHER
|
||||
// counter (it was encrypted later), so it emits after the data it trails.
|
||||
func TestMultiCoalescerRestoresTransmissionOrder(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
|
||||
// Transmission order: seq 1000 (c1), 2200 (c2), 3400 (c3).
|
||||
// Arrival order: 3400, 1000, 2200.
|
||||
if err := m.Commit(buildTCPv4(3400, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4(3400, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 1, Counter: 1}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(2200, tcpAck, pay), SortKey{Epoch: 1, Counter: 2}, testPP(buildTCPv4(2200, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 1 || len(w.writes) != 0 {
|
||||
t.Fatalf("want 1 gso write (unfragmented chain), got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
g := w.gsoWrites[0]
|
||||
if len(g.pays) != 3 {
|
||||
t.Fatalf("segs=%d want 3", len(g.pays))
|
||||
}
|
||||
const ipHdrLen = 20
|
||||
if seedSeq := binary.BigEndian.Uint32(g.hdr[ipHdrLen+4 : ipHdrLen+8]); seedSeq != 1000 {
|
||||
t.Errorf("seed seq=%d want 1000", seedSeq)
|
||||
}
|
||||
|
||||
// Retransmit: seq 1000 again but counter 4 — sorts after seq 4600 (c3).
|
||||
w.writes, w.gsoWrites, w.order = nil, nil, nil
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 1, Counter: 4}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(4600, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4(4600, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 2 {
|
||||
t.Fatalf("want 2 plain writes, got %d (gso=%d)", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
first := binary.BigEndian.Uint32(w.writes[0][24:28])
|
||||
second := binary.BigEndian.Uint32(w.writes[1][24:28])
|
||||
if first != 4600 || second != 1000 {
|
||||
t.Fatalf("emission (%d, %d), want (4600, 1000): retransmit must not overtake in-flight data", first, second)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerRestoresOrderAcrossFlows scrambles two interleaved flows;
|
||||
// the staging sort must repair each flow into one superpacket without any
|
||||
// cross-flow contamination.
|
||||
func TestMultiCoalescerRestoresOrderAcrossFlows(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
|
||||
// Transmission: A.100 (c1), B.500 (c2), A.1300 (c3), B.1700 (c4).
|
||||
// Arrival: A.1300, B.1700, A.100, B.500.
|
||||
if err := m.Commit(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay), SortKey{Epoch: 1, Counter: 3}, testPP(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay), SortKey{Epoch: 1, Counter: 4}, testPP(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay), SortKey{Epoch: 1, Counter: 1}, testPP(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay), SortKey{Epoch: 1, Counter: 2}, testPP(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 gso writes (one per flow), got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
for i, g := range w.gsoWrites {
|
||||
if len(g.pays) != 2 {
|
||||
t.Errorf("gso[%d] segs=%d want 2", i, len(g.pays))
|
||||
}
|
||||
const ipHdrLen = 20
|
||||
seedSeq := binary.BigEndian.Uint32(g.hdr[ipHdrLen+4 : ipHdrLen+8])
|
||||
sport := binary.BigEndian.Uint16(g.hdr[ipHdrLen : ipHdrLen+2])
|
||||
switch sport {
|
||||
case 1000:
|
||||
if seedSeq != 100 {
|
||||
t.Errorf("flow A seed seq=%d want 100", seedSeq)
|
||||
}
|
||||
case 3000:
|
||||
if seedSeq != 500 {
|
||||
t.Errorf("flow B seed seq=%d want 500", seedSeq)
|
||||
}
|
||||
default:
|
||||
t.Errorf("unexpected sport %d", sport)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerEpochOrdersAcrossRehandshake: a re-handshake replaces
|
||||
// the tunnel, and the replacement's counter space starts near zero — raw
|
||||
// counter order would emit the new tunnel's packets first while the old
|
||||
// tunnel's backlog is still arriving. The epoch key must dominate:
|
||||
// everything from the old tunnel emits before anything from the new one.
|
||||
func TestMultiCoalescerEpochOrdersAcrossRehandshake(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
|
||||
// New session's first data arrives before the old session's last data.
|
||||
if err := m.Commit(buildTCPv4(2200, tcpAck, pay), SortKey{Epoch: 8, Counter: 1}, testPP(buildTCPv4(2200, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 7, Counter: 9_000_000}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Same flow, contiguous seq, identical headers: after the epoch sort the
|
||||
// two segments append into one superpacket seeded by the OLD session's
|
||||
// packet.
|
||||
if len(w.gsoWrites) != 1 {
|
||||
t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
const ipHdrLen = 20
|
||||
if seedSeq := binary.BigEndian.Uint32(w.gsoWrites[0].hdr[ipHdrLen+4 : ipHdrLen+8]); seedSeq != 1000 {
|
||||
t.Errorf("seed seq=%d want 1000 (old session first)", seedSeq)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerNoUSOFallsThrough verifies that on a queue without USO
|
||||
// (older kernel: TSO but no GSO_UDP_L4) the UDP lane never comes up and UDP
|
||||
// packets still reach the kernel via verbatim rather than being lost.
|
||||
func TestMultiCoalescerNoUSOFallsThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true, noUSO: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
if m.udp != nil {
|
||||
t.Fatal("UDP lane must not come up without USO")
|
||||
}
|
||||
|
||||
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv4(1000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv4(1000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 0 {
|
||||
t.Errorf("UDP must NOT be coalesced when USO disabled, got %d gso writes", len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes) != 2 {
|
||||
t.Errorf("UDP must pass through as 2 plain writes, got %d", len(w.writes))
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerNoOffloadsStillSorts covers a queue that can't offload
|
||||
// anything. Both lane constructors refuse, so every packet rides the
|
||||
// verbatim lane — but the staging sort still applies, so emission follows
|
||||
// transmission order even without GSO.
|
||||
func TestMultiCoalescerNoOffloadsStillSorts(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: false}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
if m.tcp != nil || m.udp != nil {
|
||||
t.Fatal("no lane may come up without offloads")
|
||||
}
|
||||
pkts := [][]byte{
|
||||
buildTCPv4(1000, tcpAck, make([]byte, 1200)),
|
||||
buildUDPv4(1000, 53, make([]byte, 800)),
|
||||
buildTCPv4(2200, tcpAck, make([]byte, 1200)),
|
||||
}
|
||||
// Committed in reverse transmission order; keys carry the truth.
|
||||
for i := len(pkts) - 1; i >= 0; i-- {
|
||||
if err := m.Commit(pkts[i], SortKey{Epoch: 1, Counter: uint64(i + 1)}, testPP(pkts[i])); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 0 {
|
||||
t.Errorf("no GSO writes possible, got %d", len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes) != len(pkts) {
|
||||
t.Fatalf("want %d plain writes, got %d", len(pkts), len(w.writes))
|
||||
}
|
||||
// One lane for everything means the sorted order survives end to end.
|
||||
for i, want := range pkts {
|
||||
if !bytes.Equal(w.writes[i], want) {
|
||||
t.Errorf("write %d out of order or corrupt", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// buildUDPv6Fragment builds an IPv6 packet whose extension chain is a
|
||||
// single fragment header (NH=44) naming UDP as the terminal protocol —
|
||||
// a first fragment (offset 0, MF set) carrying the UDP header and a
|
||||
// partial payload.
|
||||
func buildUDPv6Fragment(sport, dport uint16, payload []byte) []byte {
|
||||
const ipHdrLen = 40
|
||||
const fragHdrLen = 8
|
||||
const udpHdrLen = 8
|
||||
total := ipHdrLen + fragHdrLen + udpHdrLen + len(payload)
|
||||
pkt := make([]byte, total)
|
||||
|
||||
pkt[0] = 0x60
|
||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(total-ipHdrLen))
|
||||
pkt[6] = 44 // fragment extension header
|
||||
pkt[7] = 64
|
||||
pkt[8] = 0xfe
|
||||
pkt[9] = 0x80
|
||||
pkt[23] = 1
|
||||
pkt[24] = 0xfe
|
||||
pkt[25] = 0x80
|
||||
pkt[39] = 2
|
||||
|
||||
pkt[40] = ipProtoUDP // fragment's next header
|
||||
binary.BigEndian.PutUint16(pkt[42:44], 0x0001) // offset 0, MF set
|
||||
binary.BigEndian.PutUint32(pkt[44:48], 0x1badf00) // identification
|
||||
|
||||
binary.BigEndian.PutUint16(pkt[48:50], sport)
|
||||
binary.BigEndian.PutUint16(pkt[50:52], dport)
|
||||
binary.BigEndian.PutUint16(pkt[52:54], uint16(udpHdrLen+len(payload)))
|
||||
copy(pkt[56:], payload)
|
||||
return pkt
|
||||
}
|
||||
|
||||
// TestMultiCoalescerIPv6FragmentStaysInLane locks in extension-header
|
||||
// routing: a fragment whose chain terminates in UDP must ride the UDP lane
|
||||
// as an in-lane verbatim — emitted ahead of later same-flow datagrams —
|
||||
// not the verbatim lane, which flushes after every coalescer lane and
|
||||
// would reorder it behind data that arrived after it.
|
||||
func TestMultiCoalescerIPv6FragmentStaysInLane(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
|
||||
if err := m.Commit(buildUDPv6Fragment(2000, 53, make([]byte, 512)), k.next(), testPP(buildUDPv6Fragment(2000, 53, make([]byte, 512)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 1 {
|
||||
t.Fatalf("want the fragment as 1 plain write, got %d", len(w.writes))
|
||||
}
|
||||
if len(w.gsoWrites) != 1 {
|
||||
t.Fatalf("want the two whole datagrams coalesced into 1 gso write, got %d", len(w.gsoWrites))
|
||||
}
|
||||
// Transmission order was fragment-then-data; same-lane routing must keep it.
|
||||
if w.order[0] != "write" {
|
||||
t.Fatalf("fragment must be emitted before later data (in-lane verbatim), order=%v", w.order)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerFragmentSealsUDPChains: an unparseable datagram
|
||||
// (fragment) seals every open UDP chain, so datagrams from before and after
|
||||
// it land in separate superpackets and the fragment holds its transmission-
|
||||
// order position between them.
|
||||
func TestMultiCoalescerFragmentSealsUDPChains(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6Fragment(2000, 53, make([]byte, 512)), k.next(), testPP(buildUDPv6Fragment(2000, 53, make([]byte, 512)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next(), testPP(buildUDPv6(2000, 53, make([]byte, 800)))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 gso writes (chains sealed around the fragment), got %d", len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes) != 1 {
|
||||
t.Fatalf("want the fragment as 1 plain write, got %d", len(w.writes))
|
||||
}
|
||||
want := []string{"gso", "write", "gso"}
|
||||
if len(w.order) != 3 || w.order[0] != want[0] || w.order[1] != want[1] || w.order[2] != want[2] {
|
||||
t.Fatalf("emission order = %v, want %v", w.order, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiCoalescerNoTSOFallsThrough mirrors the no-TSO case.
|
||||
func TestMultiCoalescerNoTSOFallsThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true, noTSO: true}
|
||||
m := newTestMultiCoalescer(t, w)
|
||||
k := &keySeq{epoch: 1}
|
||||
if m.tcp != nil {
|
||||
t.Fatal("TCP lane must not come up without TSO")
|
||||
}
|
||||
|
||||
pay := make([]byte, 1200)
|
||||
if err := m.Commit(buildTCPv4(1000, tcpAck, pay), k.next(), testPP(buildTCPv4(1000, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Commit(buildTCPv4(2200, tcpAck, pay), k.next(), testPP(buildTCPv4(2200, tcpAck, pay))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 0 {
|
||||
t.Errorf("TCP must NOT be coalesced when TSO disabled, got %d gso writes", len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes) != 2 {
|
||||
t.Errorf("TCP must pass through as 2 plain writes, got %d", len(w.writes))
|
||||
}
|
||||
}
|
||||
|
||||
// testPP derives the ParsedPacket newPacket would produce for the packet
|
||||
// shapes the tests build: plain v4/v6, v4 with options or fragment bits set,
|
||||
// and the single-fragment-header v6 shape from buildUDPv6Fragment. Anything
|
||||
// unrecognizable stays zero (proto 0 routes to the passthrough lane).
|
||||
func testPP(pkt []byte) *firewall.ParsedPacket {
|
||||
pp := &firewall.ParsedPacket{}
|
||||
if len(pkt) < 20 {
|
||||
return pp
|
||||
}
|
||||
switch pkt[0] >> 4 {
|
||||
case 4:
|
||||
pp.Protocol = pkt[9]
|
||||
pp.IPHdrLen = int(pkt[0]&0x0f) * 4
|
||||
pp.FragAny = binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0
|
||||
case 6:
|
||||
pp.Protocol = pkt[6]
|
||||
pp.IPHdrLen = 40
|
||||
if pp.Protocol == 44 { // fragment extension header
|
||||
pp.Protocol = pkt[40]
|
||||
pp.IPHdrLen = 48
|
||||
pp.FragAny = true
|
||||
}
|
||||
}
|
||||
return pp
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"io"
|
||||
)
|
||||
|
||||
// Passthrough is MultiCoalescer's verbatim lane: no batching, packets are written at Flush in the
|
||||
// order enqueued.
|
||||
type Passthrough struct {
|
||||
out io.Writer
|
||||
slots [][]byte
|
||||
}
|
||||
|
||||
func NewPassthrough(w io.Writer) *Passthrough {
|
||||
return &Passthrough{
|
||||
out: w,
|
||||
slots: make([][]byte, 0, 128),
|
||||
}
|
||||
}
|
||||
|
||||
// enqueue accepts one packet, already sorted into transmission order by dispatch.
|
||||
func (p *Passthrough) enqueue(pkt []byte) error {
|
||||
p.slots = append(p.slots, pkt)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *Passthrough) Flush() error {
|
||||
var firstErr error
|
||||
for _, s := range p.slots {
|
||||
_, err := p.out.Write(s)
|
||||
if err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
clear(p.slots)
|
||||
p.slots = p.slots[:0]
|
||||
return firstErr
|
||||
}
|
||||
@@ -0,0 +1,472 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"log/slog"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
)
|
||||
|
||||
// ipProtoTCP is the IANA protocol number for TCP. Defined here to help Windows out.
|
||||
const ipProtoTCP = 6
|
||||
|
||||
// tcpCoalesceBufSize caps total bytes per superpacket. Mirrors the kernel's
|
||||
// sk_gso_max_size of ~64KiB; anything beyond this would be rejected anyway.
|
||||
const tcpCoalesceBufSize = 65535
|
||||
|
||||
// tcpCoalesceMaxSegs caps how many segments we'll coalesce into a single
|
||||
// superpacket. Keeping this well below the kernel's TSO ceiling bounds latency.
|
||||
const tcpCoalesceMaxSegs = 64
|
||||
|
||||
// coalesceSlot is one entry in the coalescer's ordered event queue. A verbatim slot holds a single
|
||||
// borrowed packet emitted as-is (pure ACK, non-admissible TCP, unparseable, or oversize seed); a
|
||||
// non-verbatim slot is an in-progress coalesced superpacket. payIovs are borrowed slices of the
|
||||
// caller's plaintext buffers; the caller must keep them alive until Flush.
|
||||
type coalesceSlot struct {
|
||||
verbatim bool
|
||||
// rawPkt is borrowed: the whole packet for verbatim slots, the seed packet for coalesce
|
||||
// slots. A slot that never grows past one segment is emitted from rawPkt so its original
|
||||
// (already valid) L4 checksum ships DATA_VALID instead of making the kernel recompute it.
|
||||
// A multi-segment slot's superpacket header is rawPkt's, patched in place at flush.
|
||||
rawPkt []byte
|
||||
|
||||
fk flowKey
|
||||
hdrLen int
|
||||
ipHdrLen int
|
||||
isV6 bool
|
||||
gsoSize int
|
||||
numSeg int
|
||||
totalPay int
|
||||
nextSeq uint32
|
||||
payIovs [][]byte
|
||||
}
|
||||
|
||||
// TCPCoalescer accumulates adjacent in-flow TCP data segments across multiple concurrent flows and
|
||||
// emits each flow's run as a single TSO superpacket via tio.GSOWriter. Input must be in sender
|
||||
// transmission order (MultiCoalescer sorts by (epoch, counter) before dispatch); slots are emitted
|
||||
// in creation order, so emission reproduces transmission order except for the pure-ACK case in
|
||||
// commitParsed. Owns no locks; one coalescer per TUN write queue.
|
||||
type TCPCoalescer struct {
|
||||
w tio.GSOWriter
|
||||
|
||||
// slots is the ordered event queue. Flush walks it once and emits each
|
||||
// entry as either a WriteGSO (coalesced) or a w.Write (verbatim).
|
||||
slots []*coalesceSlot
|
||||
// openSlots maps a flow key to its open slot so new segments can extend an in-progress
|
||||
// superpacket in O(1). Removal is what closes a chain: on PSH or a short last segment, on a
|
||||
// non-admissible packet for the flow, or in Flush.
|
||||
openSlots map[flowKey]*coalesceSlot
|
||||
// lastSlot caches the most recently touched open slot. Bulk traffic
|
||||
// arrives in same-flow runs (single-flow steady state, or GRO bursts
|
||||
// under multi-flow), so comparing the incoming key against the cached
|
||||
// slot's own fk lets the hot path skip the map lookup (and the aeshash
|
||||
// of a 38-byte key) for the length of each run.
|
||||
// Kept in lockstep with openSlots: nil whenever the slot it pointed
|
||||
// at is removed.
|
||||
lastSlot *coalesceSlot
|
||||
pool []*coalesceSlot // free list for reuse
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
// NewTCPCoalescer wraps w, returning nil if w can't accept GSO_TCP writes.
|
||||
func NewTCPCoalescer(w io.Writer, l *slog.Logger) *TCPCoalescer {
|
||||
gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
return &TCPCoalescer{
|
||||
w: gw,
|
||||
slots: make([]*coalesceSlot, 0, initialSlots),
|
||||
openSlots: make(map[flowKey]*coalesceSlot, initialSlots),
|
||||
pool: make([]*coalesceSlot, 0, initialSlots),
|
||||
l: l,
|
||||
}
|
||||
}
|
||||
|
||||
// parsedTCP holds the fields extracted from a single parse so later steps
|
||||
// (admission, slot lookup, canAppend) don't re-walk the header.
|
||||
type parsedTCP struct {
|
||||
fk flowKey
|
||||
ipHdrLen int
|
||||
hdrLen int
|
||||
payLen int
|
||||
seq uint32
|
||||
flags byte
|
||||
}
|
||||
|
||||
// parseAt extracts the flow key and IP/TCP offsets for a packet the dispatcher already knows is
|
||||
// TCP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on
|
||||
// entry and is filled in place; see flowKey.parseIPAt for why. Returns false for malformed input
|
||||
// or any shape that must not coalesce (IPv4 options/fragmentation, IPv6 extension headers).
|
||||
func (p *parsedTCP) parseAt(pkt []byte, ipHdrLen int) bool {
|
||||
trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return p.parseTail(trimmed, ipHdrLen)
|
||||
}
|
||||
|
||||
// parseTail layers the TCP-header parse on a validated IP prologue. pkt is the trimmed packet;
|
||||
// fk's addresses are already filled.
|
||||
func (p *parsedTCP) parseTail(pkt []byte, ipHdrLen int) bool {
|
||||
if len(pkt) < ipHdrLen+20 {
|
||||
return false
|
||||
}
|
||||
tcpOff := int(pkt[ipHdrLen+12]>>4) * 4
|
||||
if tcpOff < 20 || tcpOff > 60 {
|
||||
return false
|
||||
}
|
||||
if len(pkt) < ipHdrLen+tcpOff {
|
||||
return false
|
||||
}
|
||||
p.ipHdrLen = ipHdrLen
|
||||
p.hdrLen = ipHdrLen + tcpOff
|
||||
p.payLen = len(pkt) - p.hdrLen
|
||||
p.fk.sport = binary.BigEndian.Uint16(pkt[ipHdrLen : ipHdrLen+2])
|
||||
p.fk.dport = binary.BigEndian.Uint16(pkt[ipHdrLen+2 : ipHdrLen+4])
|
||||
p.seq = binary.BigEndian.Uint32(pkt[ipHdrLen+4 : ipHdrLen+8])
|
||||
p.flags = pkt[ipHdrLen+13]
|
||||
return true
|
||||
}
|
||||
|
||||
// TCP flag bits (byte 13 of the TCP header). Only the bits the coalescer consults are named;
|
||||
// FIN/SYN/RST/URG/CWR are rejected by the negative mask in commitParsed.
|
||||
const (
|
||||
tcpFlagPsh = 0x08
|
||||
tcpFlagAck = 0x10
|
||||
tcpFlagEce = 0x40
|
||||
)
|
||||
|
||||
// sealAllOpen closes every open coalesce chain. Called for unparseable packets: the flow key is
|
||||
// unknown, so any open chain could otherwise absorb later data and emit it ahead of this packet.
|
||||
func (c *TCPCoalescer) sealAllOpen() {
|
||||
clear(c.openSlots)
|
||||
c.lastSlot = nil
|
||||
}
|
||||
|
||||
// sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips
|
||||
// hashing the 38-byte key when no chains are open (e.g. ack-dominant queues).
|
||||
func (c *TCPCoalescer) sealFlow(fk flowKey) {
|
||||
if len(c.openSlots) == 0 {
|
||||
return
|
||||
}
|
||||
if last := c.lastSlot; last != nil && last.fk == fk {
|
||||
c.lastSlot = nil
|
||||
}
|
||||
delete(c.openSlots, fk)
|
||||
}
|
||||
|
||||
// commitStaged commits one staged packet dispatch routed to this lane. A shape the lane cannot
|
||||
// coalesce (any fragmentation, unparseable header) seals every open chain
|
||||
// and rides the lane as an in-lane verbatim, still in transmission order.
|
||||
func (c *TCPCoalescer) commitStaged(sp stagedPacket) error {
|
||||
if sp.fragAny {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(sp.pkt)
|
||||
return nil
|
||||
}
|
||||
var info parsedTCP
|
||||
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(sp.pkt)
|
||||
return nil
|
||||
}
|
||||
return c.commitParsed(sp.pkt, &info)
|
||||
}
|
||||
|
||||
// commitParsed commits one parsed TCP packet. The caller (dispatch, via parseAt) supplies a
|
||||
// valid parse so the header is not re-walked here.
|
||||
func (c *TCPCoalescer) commitParsed(pkt []byte, info *parsedTCP) error {
|
||||
// Admission: only ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE may ride a coalesce chain. CWR marks a
|
||||
// one-shot congestion transition the receiver must observe at a segment boundary. NB: AccECN
|
||||
// reuses CWR as ACE counter bits; revisit this check if inner hosts adopt AccECN.
|
||||
if info.flags&tcpFlagAck == 0 || info.flags&^(tcpFlagAck|tcpFlagPsh|tcpFlagEce) != 0 {
|
||||
// SYN/FIN/RST/URG/CWR must be observed in sequence. Seal the flow's open slot so later
|
||||
// in-flow packets cannot extend it and emit ahead of this verbatim.
|
||||
c.sealFlow(info.fk)
|
||||
c.addVerbatim(pkt)
|
||||
return nil
|
||||
}
|
||||
if info.payLen == 0 {
|
||||
// Pure ACK: no ordering obligation toward the flow's data. Delivering it after
|
||||
// later-transmitted data only makes it a stale ACK, which receivers ignore. Not sealing
|
||||
// keeps a bidirectional flow's data run coalescing across interleaved peer ACKs, matching
|
||||
// kernel GRO. This is the only place emission deviates from transmission order.
|
||||
c.addVerbatim(pkt)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Cached-slot fast path. Arrival isn't per-packet interleaved even with
|
||||
// many flows: wire-side GRO delivers runs of same-flow packets
|
||||
// (deliverSegments splits a superdatagram into up to 64), so the cache
|
||||
// hits for the length of each run and a miss costs one fk compare
|
||||
// before the map lookup carries the weight.
|
||||
var open *coalesceSlot
|
||||
if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||
open = last
|
||||
} else {
|
||||
open = c.openSlots[info.fk]
|
||||
}
|
||||
if open != nil {
|
||||
if c.canAppend(open, pkt, info) {
|
||||
if c.appendPayload(open, pkt, info) {
|
||||
// Chain closed (PSH or short segment): stop extending it.
|
||||
c.sealFlow(info.fk)
|
||||
} else {
|
||||
c.lastSlot = open
|
||||
}
|
||||
return nil
|
||||
}
|
||||
// Can't extend (seq gap from upstream loss, header change, or a full
|
||||
// chain): evict it from openSlots and fall through to seed a fresh slot.
|
||||
c.sealFlow(info.fk)
|
||||
}
|
||||
c.seed(pkt, info)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) Flush() error {
|
||||
var first error
|
||||
for _, s := range c.slots {
|
||||
var err error
|
||||
if s.verbatim || s.numSeg == 1 {
|
||||
// A slot that never grew is byte-identical to its seed packet; ship the original so
|
||||
// its valid checksum rides the DATA_VALID path instead of a kernel software csum.
|
||||
// rawPkt is only mutated once numSeg >= 2 (PSH propagate, flush patches), so it is
|
||||
// pristine here.
|
||||
_, err = c.w.Write(s.rawPkt)
|
||||
} else {
|
||||
err = c.flushSlot(s)
|
||||
}
|
||||
if err != nil && first == nil {
|
||||
first = err
|
||||
}
|
||||
c.release(s)
|
||||
}
|
||||
clear(c.slots)
|
||||
c.slots = c.slots[:0]
|
||||
clear(c.openSlots)
|
||||
c.lastSlot = nil
|
||||
|
||||
return first
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) addVerbatim(pkt []byte) {
|
||||
s := c.take()
|
||||
s.verbatim = true
|
||||
s.rawPkt = pkt
|
||||
c.slots = append(c.slots, s)
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) seed(pkt []byte, info *parsedTCP) {
|
||||
if info.hdrLen+info.payLen > tcpCoalesceBufSize {
|
||||
// Pathological shape that can't ride a superpacket; emit as-is. No chain for this flow can
|
||||
// be open here (commitParsed evicts before seeding), so sealFlow is defense in depth
|
||||
// against a stale cache entry absorbing later data.
|
||||
c.sealFlow(info.fk)
|
||||
c.addVerbatim(pkt)
|
||||
return
|
||||
}
|
||||
s := c.take()
|
||||
s.verbatim = false
|
||||
// rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is
|
||||
// the superpacket header flushSlot patches in place.
|
||||
s.rawPkt = pkt
|
||||
s.hdrLen = info.hdrLen
|
||||
s.ipHdrLen = info.ipHdrLen
|
||||
s.isV6 = info.fk.isV6
|
||||
s.fk = info.fk
|
||||
s.gsoSize = info.payLen
|
||||
s.numSeg = 1
|
||||
s.totalPay = info.payLen
|
||||
s.nextSeq = info.seq + uint32(info.payLen)
|
||||
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||
c.slots = append(c.slots, s)
|
||||
if info.flags&tcpFlagPsh == 0 {
|
||||
c.openSlots[info.fk] = s
|
||||
c.lastSlot = s
|
||||
} else {
|
||||
// PSH on the seed closes the chain immediately; it is never registered as open.
|
||||
// Drop any stale entry for this flow too (defense in depth, unreachable if lastSlot's lockstep invariant holds).
|
||||
c.sealFlow(info.fk)
|
||||
}
|
||||
}
|
||||
|
||||
// canAppend reports whether info's packet extends the slot's seed: same header shape and stable
|
||||
// contents, adjacent seq, not oversized. A closed chain never reaches here; closing removes the
|
||||
// slot from openSlots, the only path in. The header fields read from rawPkt are always pristine:
|
||||
// the only pre-flush mutation is the PSH propagate, which also closes the chain.
|
||||
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info *parsedTCP) bool {
|
||||
if info.hdrLen != s.hdrLen {
|
||||
return false
|
||||
}
|
||||
if info.seq != s.nextSeq {
|
||||
return false
|
||||
}
|
||||
if s.numSeg >= tcpCoalesceMaxSegs {
|
||||
return false
|
||||
}
|
||||
if info.payLen > s.gsoSize {
|
||||
return false
|
||||
}
|
||||
if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize {
|
||||
return false
|
||||
}
|
||||
// ECE state must be stable across a burst.
|
||||
// Receivers expect the flag set on every segment of a CE-echoing window or none.
|
||||
seedFlags := s.rawPkt[s.ipHdrLen+13]
|
||||
if (seedFlags^info.flags)&tcpFlagEce != 0 {
|
||||
return false
|
||||
}
|
||||
if !s.isV6 && !ipv4CanCoalesceID(s.rawPkt, pkt, s.numSeg) {
|
||||
return false
|
||||
}
|
||||
if !headersMatch(s.rawPkt[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// appendPayload folds info's packet into s and reports whether the chain is now closed: the
|
||||
// segment was sub-gsoSize (kernel TSO allows only the final segment to be short) or carried PSH.
|
||||
// The caller must deregister a closed slot from openSlots.
|
||||
func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info *parsedTCP) bool {
|
||||
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||
s.numSeg++
|
||||
s.totalPay += info.payLen
|
||||
s.nextSeq = info.seq + uint32(info.payLen)
|
||||
if info.flags&tcpFlagPsh != 0 {
|
||||
// Propagate PSH into the seed header so kernel TSO sets it on the last segment. Mutating
|
||||
// rawPkt is safe: PSH also closes the chain, so no admission check re-reads this header.
|
||||
s.rawPkt[s.ipHdrLen+13] |= tcpFlagPsh
|
||||
}
|
||||
return info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) take() *coalesceSlot {
|
||||
if n := len(c.pool); n > 0 {
|
||||
s := c.pool[n-1]
|
||||
c.pool[n-1] = nil
|
||||
c.pool = c.pool[:n-1]
|
||||
return s
|
||||
}
|
||||
return &coalesceSlot{}
|
||||
}
|
||||
|
||||
func (c *TCPCoalescer) release(s *coalesceSlot) {
|
||||
clear(s.payIovs)
|
||||
*s = coalesceSlot{payIovs: s.payIovs[:0]}
|
||||
c.pool = append(c.pool, s)
|
||||
}
|
||||
|
||||
// flushSlot patches the superpacket header in place in rawPkt (total length, IPv4 header
|
||||
// checksum, pseudo-header checksum seed) and calls WriteGSO. The slot is released right after,
|
||||
// so nothing re-reads the patched header. Does not remove the slot from c.slots.
|
||||
func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
|
||||
total := s.hdrLen + s.totalPay
|
||||
l4Len := total - s.ipHdrLen
|
||||
hdr := s.rawPkt[:s.hdrLen]
|
||||
|
||||
if s.isV6 {
|
||||
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
|
||||
} else {
|
||||
binary.BigEndian.PutUint16(hdr[2:4], uint16(total))
|
||||
hdr[10] = 0
|
||||
hdr[11] = 0
|
||||
binary.BigEndian.PutUint16(hdr[10:12], ipv4HdrChecksum(hdr[:s.ipHdrLen]))
|
||||
}
|
||||
|
||||
var psum uint32
|
||||
if s.isV6 {
|
||||
psum = pseudoSumIPv6(hdr[8:24], hdr[24:40], ipProtoTCP, l4Len)
|
||||
} else {
|
||||
psum = pseudoSumIPv4(hdr[12:16], hdr[16:20], ipProtoTCP, l4Len)
|
||||
}
|
||||
tcsum := s.ipHdrLen + 16
|
||||
binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum))
|
||||
|
||||
return c.w.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoTCP)
|
||||
}
|
||||
|
||||
// headersMatch compares two IP+TCP header prefixes for byte-for-byte
|
||||
// equality on every field that must be identical across coalesced
|
||||
// segments. Size/IPID/IPCsum/seq/flags/tcpCsum are masked out.
|
||||
func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
if !ipHeadersMatch(a, b, isV6) {
|
||||
return false
|
||||
}
|
||||
// TCP: compare [0:4] ports, [8:13] ack+dataoff, [14:16] window,
|
||||
// [18:tcpHdrLen] options (incl. urgent).
|
||||
tcp := ipHdrLen
|
||||
if !bytes.Equal(a[tcp:tcp+4], b[tcp:tcp+4]) {
|
||||
return false
|
||||
}
|
||||
if !bytes.Equal(a[tcp+8:tcp+13], b[tcp+8:tcp+13]) {
|
||||
return false
|
||||
}
|
||||
if !bytes.Equal(a[tcp+14:tcp+16], b[tcp+14:tcp+16]) {
|
||||
return false
|
||||
}
|
||||
if !bytes.Equal(a[tcp+18:], b[tcp+18:]) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must
|
||||
// already have its checksum field zeroed) and returns the folded/inverted
|
||||
// 16-bit value to store.
|
||||
func ipv4HdrChecksum(hdr []byte) uint16 {
|
||||
var sum uint32
|
||||
for i := 0; i+1 < len(hdr); i += 2 {
|
||||
sum += uint32(binary.BigEndian.Uint16(hdr[i : i+2]))
|
||||
}
|
||||
if len(hdr)%2 == 1 {
|
||||
sum += uint32(hdr[len(hdr)-1]) << 8
|
||||
}
|
||||
for sum>>16 != 0 {
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
}
|
||||
return ^uint16(sum)
|
||||
}
|
||||
|
||||
// pseudoSumIPv4 / pseudoSumIPv6 build the L4 pseudo-header partial sum
|
||||
// expected by the virtio NEEDS_CSUM kernel path: the 32-bit accumulator
|
||||
// before folding. proto selects the L4 (TCP or UDP); the UDP coalescer
|
||||
// reuses these helpers.
|
||||
func pseudoSumIPv4(src, dst []byte, proto byte, l4Len int) uint32 {
|
||||
var sum uint32
|
||||
sum += uint32(binary.BigEndian.Uint16(src[0:2]))
|
||||
sum += uint32(binary.BigEndian.Uint16(src[2:4]))
|
||||
sum += uint32(binary.BigEndian.Uint16(dst[0:2]))
|
||||
sum += uint32(binary.BigEndian.Uint16(dst[2:4]))
|
||||
sum += uint32(proto)
|
||||
sum += uint32(l4Len)
|
||||
return sum
|
||||
}
|
||||
|
||||
func pseudoSumIPv6(src, dst []byte, proto byte, l4Len int) uint32 {
|
||||
var sum uint32
|
||||
for i := 0; i < 16; i += 2 {
|
||||
sum += uint32(binary.BigEndian.Uint16(src[i : i+2]))
|
||||
sum += uint32(binary.BigEndian.Uint16(dst[i : i+2]))
|
||||
}
|
||||
sum += uint32(l4Len >> 16)
|
||||
sum += uint32(l4Len & 0xffff)
|
||||
sum += uint32(proto)
|
||||
return sum
|
||||
}
|
||||
|
||||
// foldOnceNoInvert folds the 32-bit accumulator to 16 bits and returns it unchanged (no one's complement).
|
||||
// This is what virtio NEEDS_CSUM wants in the L4 checksum field
|
||||
func foldOnceNoInvert(sum uint32) uint16 {
|
||||
for sum>>16 != 0 {
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
}
|
||||
return uint16(sum)
|
||||
}
|
||||
@@ -0,0 +1,214 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/test"
|
||||
)
|
||||
|
||||
// nopTunWriter is a zero-alloc tio.GSOWriter for benchmarks. Discards
|
||||
// everything but satisfies the interface the coalescer detects.
|
||||
type nopTunWriter struct{}
|
||||
|
||||
func (nopTunWriter) Write(p []byte) (int, error) { return len(p), nil }
|
||||
func (nopTunWriter) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, _ tio.GSOProto) error {
|
||||
return nil
|
||||
}
|
||||
func (nopTunWriter) Capabilities() tio.Capabilities {
|
||||
return tio.Capabilities{TSO: true, USO: true}
|
||||
}
|
||||
|
||||
// buildTCPv4BulkFlow returns a slice of N adjacent ACK-only TCP segments
|
||||
// on a single 5-tuple, each carrying payloadLen bytes. Seq numbers are
|
||||
// contiguous so every packet is coalesceable onto the previous one.
|
||||
func buildTCPv4BulkFlow(n, payloadLen int) [][]byte {
|
||||
pkts := make([][]byte, n)
|
||||
pay := make([]byte, payloadLen)
|
||||
seq := uint32(1000)
|
||||
for i := range n {
|
||||
pkts[i] = buildTCPv4(seq, tcpAck, pay)
|
||||
seq += uint32(payloadLen)
|
||||
}
|
||||
return pkts
|
||||
}
|
||||
|
||||
// buildTCPv4Interleaved returns nFlows * perFlow packets with per-flow
|
||||
// seq continuity but round-robin across flows — worst case for any
|
||||
// "last-slot" cache.
|
||||
func buildTCPv4Interleaved(nFlows, perFlow, payloadLen int) [][]byte {
|
||||
pay := make([]byte, payloadLen)
|
||||
seqs := make([]uint32, nFlows)
|
||||
for i := range seqs {
|
||||
seqs[i] = uint32(1000 + i*1000000)
|
||||
}
|
||||
pkts := make([][]byte, 0, nFlows*perFlow)
|
||||
for range perFlow {
|
||||
for f := range nFlows {
|
||||
sport := uint16(10000 + f)
|
||||
pkts = append(pkts, buildTCPv4Ports(sport, 2000, seqs[f], tcpAck, pay))
|
||||
seqs[f] += uint32(payloadLen)
|
||||
}
|
||||
}
|
||||
return pkts
|
||||
}
|
||||
|
||||
// buildTCPv4RunInterleaved returns nFlows*perFlow packets delivered in
|
||||
// runs of runLen per flow — the arrival pattern wire-side GRO actually
|
||||
// produces (deliverSegments splits each superdatagram into up to 64
|
||||
// same-flow packets back to back). Contrast with buildTCPv4Interleaved's
|
||||
// per-packet round-robin, the adversarial worst case for a last-slot cache.
|
||||
func buildTCPv4RunInterleaved(nFlows, perFlow, runLen, payloadLen int) [][]byte {
|
||||
pay := make([]byte, payloadLen)
|
||||
seqs := make([]uint32, nFlows)
|
||||
for i := range seqs {
|
||||
seqs[i] = uint32(1000 + i*1000000)
|
||||
}
|
||||
pkts := make([][]byte, 0, nFlows*perFlow)
|
||||
for done := 0; done < perFlow; done += runLen {
|
||||
for f := range nFlows {
|
||||
sport := uint16(10000 + f)
|
||||
for range runLen {
|
||||
pkts = append(pkts, buildTCPv4Ports(sport, 2000, seqs[f], tcpAck, pay))
|
||||
seqs[f] += uint32(payloadLen)
|
||||
}
|
||||
}
|
||||
}
|
||||
return pkts
|
||||
}
|
||||
|
||||
// buildICMPv4 returns a minimal non-TCP packet that takes the verbatim
|
||||
// branch in Commit.
|
||||
func buildICMPv4() []byte {
|
||||
pkt := make([]byte, 28)
|
||||
pkt[0] = 0x45
|
||||
binary.BigEndian.PutUint16(pkt[2:4], 28)
|
||||
pkt[9] = 1 // ICMP
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||
return pkt
|
||||
}
|
||||
|
||||
// runCommitBench drives Commit over pkts batchSize at a time, flushing
|
||||
// between batches, and reports per-packet cost.
|
||||
func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||
b.Helper()
|
||||
c := newTestTCPCoalescer(b, nopTunWriter{})
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkts[0])))
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
pkt := pkts[i%len(pkts)]
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if (i+1)%batchSize == 0 {
|
||||
if err := c.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
// Drain any trailing partial batch so slot state doesn't leak across runs.
|
||||
_ = c.Flush()
|
||||
}
|
||||
|
||||
// BenchmarkCommitSingleFlow is the bulk-TCP steady state: one flow,
|
||||
// contiguous seq, 1200-byte payloads. Every packet past the seed should
|
||||
// append onto the open slot. This is the case we most care about.
|
||||
func BenchmarkCommitSingleFlow(b *testing.B) {
|
||||
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
|
||||
runCommitBench(b, pkts, tcpCoalesceMaxSegs)
|
||||
}
|
||||
|
||||
// BenchmarkCommitInterleaved4 has 4 concurrent bulk flows round-robined.
|
||||
// A single-entry fast-path cache will miss on every packet; an N-way
|
||||
// cache or map lookup carries the weight.
|
||||
func BenchmarkCommitInterleaved4(b *testing.B) {
|
||||
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
|
||||
runCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkCommitInterleaved16 stresses the map at higher flow counts.
|
||||
func BenchmarkCommitInterleaved16(b *testing.B) {
|
||||
pkts := buildTCPv4Interleaved(16, tcpCoalesceMaxSegs, 1200)
|
||||
runCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkCommitRunInterleaved4 is 4 concurrent flows arriving in
|
||||
// GRO-burst runs of 16 — the realistic multi-flow pattern. A last-slot
|
||||
// cache hits for the length of each run; the per-packet round-robin
|
||||
// benches above are its worst case.
|
||||
func BenchmarkCommitRunInterleaved4(b *testing.B) {
|
||||
pkts := buildTCPv4RunInterleaved(4, tcpCoalesceMaxSegs, 16, 1200)
|
||||
runCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkCommitPassthrough exercises the non-TCP branch: parseBase
|
||||
// bails early and addVerbatim is the only work.
|
||||
func BenchmarkCommitPassthrough(b *testing.B) {
|
||||
pkt := buildICMPv4()
|
||||
pkts := make([][]byte, 64)
|
||||
for i := range pkts {
|
||||
pkts[i] = pkt
|
||||
}
|
||||
runCommitBench(b, pkts, 64)
|
||||
}
|
||||
|
||||
// BenchmarkCommitNonCoalesceableTCP sends SYN|ACK packets on one flow.
|
||||
// Each packet takes the "TCP but not admissible" branch which does a
|
||||
// map delete + verbatim. Measures the seal-without-slot cost.
|
||||
func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
|
||||
pay := make([]byte, 0)
|
||||
pkts := make([][]byte, 64)
|
||||
for i := range pkts {
|
||||
pkts[i] = buildTCPv4(uint32(1000+i), tcpSyn|tcpAck, pay)
|
||||
}
|
||||
runCommitBench(b, pkts, 64)
|
||||
}
|
||||
|
||||
// runMultiCommitBench drives MultiCoalescer.Commit with in-order keys, so
|
||||
// it includes the staging sort's already-sorted fast path plus the
|
||||
// dispatch-time parse — the full steady-state cost of the batcher. The
|
||||
// ParsedPackets are precomputed: in production they fall out of the
|
||||
// firewall's newPacket, which this bench does not model.
|
||||
func runMultiCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||
b.Helper()
|
||||
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger())
|
||||
pps := make([]*firewall.ParsedPacket, len(pkts))
|
||||
for i, p := range pkts {
|
||||
pps[i] = testPP(p)
|
||||
}
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkts[0])))
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
j := i % len(pkts)
|
||||
if err := m.Commit(pkts[j], SortKey{Epoch: 1, Counter: uint64(i + 1)}, pps[j]); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if (i+1)%batchSize == 0 {
|
||||
if err := m.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
_ = m.Flush()
|
||||
}
|
||||
|
||||
// BenchmarkMultiCommitSingleFlow is the multi-lane analogue of
|
||||
// BenchmarkCommitSingleFlow — same workload but routed through the
|
||||
// dispatcher. The delta vs the single-lane bench measures dispatcher
|
||||
// overhead.
|
||||
func BenchmarkMultiCommitSingleFlow(b *testing.B) {
|
||||
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
|
||||
runMultiCommitBench(b, pkts, tcpCoalesceMaxSegs)
|
||||
}
|
||||
|
||||
// BenchmarkMultiCommitInterleaved4 mirrors BenchmarkCommitInterleaved4
|
||||
// through the dispatcher.
|
||||
func BenchmarkMultiCommitInterleaved4(b *testing.B) {
|
||||
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
|
||||
runMultiCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,59 @@
|
||||
package batch
|
||||
|
||||
import "net/netip"
|
||||
|
||||
const SendBatchCap = 128
|
||||
|
||||
// batchWriter is the minimal subset of udp.Conn needed by SendBatch to flush.
|
||||
type batchWriter interface {
|
||||
WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error)
|
||||
}
|
||||
|
||||
// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch.
|
||||
// One SendBatch is owned by each listenIn goroutine; no locking is needed.
|
||||
// Slots are backed by an Arena (see its docs)
|
||||
type SendBatch struct {
|
||||
out batchWriter
|
||||
bufs [][]byte
|
||||
dsts []netip.AddrPort
|
||||
arena *Arena
|
||||
}
|
||||
|
||||
// NewSendBatch makes a SendBatch with batchCap slots and an arenaSize byte buffer for slices to back those slots
|
||||
func NewSendBatch(out batchWriter, batchCap, arenaSize int) *SendBatch {
|
||||
return &SendBatch{
|
||||
out: out,
|
||||
bufs: make([][]byte, 0, batchCap),
|
||||
dsts: make([]netip.AddrPort, 0, batchCap),
|
||||
arena: NewArena(arenaSize),
|
||||
}
|
||||
}
|
||||
|
||||
func (b *SendBatch) Reserve(sz int) []byte {
|
||||
return b.arena.Reserve(sz)
|
||||
}
|
||||
|
||||
// Len reports how many packets are queued for the next Flush. Callers use
|
||||
// it to flush incrementally once a full sendmmsg batch has accumulated,
|
||||
// bounding how long the first packet of a large read batch waits.
|
||||
func (b *SendBatch) Len() int { return len(b.bufs) }
|
||||
|
||||
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort) {
|
||||
b.bufs = append(b.bufs, pkt)
|
||||
b.dsts = append(b.dsts, dst)
|
||||
}
|
||||
|
||||
// Flush writes every queued packet and reports how many actually went out. A short count means some destinations
|
||||
// were undeliverable; the batch is drained either way.
|
||||
func (b *SendBatch) Flush() (int, error) {
|
||||
var err error
|
||||
written := 0
|
||||
if len(b.bufs) > 0 {
|
||||
written, err = b.out.WriteBatch(b.bufs, b.dsts)
|
||||
}
|
||||
clear(b.bufs)
|
||||
b.bufs = b.bufs[:0]
|
||||
b.dsts = b.dsts[:0]
|
||||
b.arena.Reset()
|
||||
return written, err
|
||||
}
|
||||
@@ -0,0 +1,122 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type fakeBatchWriter struct {
|
||||
bufs [][]byte
|
||||
addrs []netip.AddrPort
|
||||
}
|
||||
|
||||
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
|
||||
// Snapshot — SendBatch.Flush nils its slot pointers right after WriteBatch
|
||||
// returns, so tests must capture data before that happens.
|
||||
w.bufs = make([][]byte, len(bufs))
|
||||
for i, b := range bufs {
|
||||
cp := make([]byte, len(b))
|
||||
copy(cp, b)
|
||||
w.bufs[i] = cp
|
||||
}
|
||||
w.addrs = append(w.addrs[:0], addrs...)
|
||||
return len(bufs), nil
|
||||
}
|
||||
|
||||
func TestSendBatchReserveCommitFlush(t *testing.T) {
|
||||
fw := &fakeBatchWriter{}
|
||||
b := NewSendBatch(fw, 4, 32)
|
||||
|
||||
ap := netip.MustParseAddrPort("10.0.0.1:4242")
|
||||
for i := 0; i < 4; i++ {
|
||||
slot := b.Reserve(32)
|
||||
if cap(slot) != 32 {
|
||||
t.Fatalf("slot %d: cap=%d want 32", i, cap(slot))
|
||||
}
|
||||
pkt := append(slot[:0], byte(i), byte(i+1), byte(i+2))
|
||||
b.Commit(pkt, ap)
|
||||
}
|
||||
if _, err := b.Flush(); err != nil {
|
||||
t.Fatalf("Flush: %v", err)
|
||||
}
|
||||
if len(fw.bufs) != 4 {
|
||||
t.Fatalf("WriteBatch got %d bufs want 4", len(fw.bufs))
|
||||
}
|
||||
for i, buf := range fw.bufs {
|
||||
if len(buf) != 3 || buf[0] != byte(i) {
|
||||
t.Errorf("buf %d: %x", i, buf)
|
||||
}
|
||||
if fw.addrs[i] != ap {
|
||||
t.Errorf("addr %d: got %v want %v", i, fw.addrs[i], ap)
|
||||
}
|
||||
}
|
||||
|
||||
// Flush again with nothing committed — should be a no-op.
|
||||
fw.bufs = nil
|
||||
if _, err := b.Flush(); err != nil {
|
||||
t.Fatalf("empty Flush: %v", err)
|
||||
}
|
||||
if fw.bufs != nil {
|
||||
t.Fatalf("empty Flush triggered WriteBatch")
|
||||
}
|
||||
|
||||
// Reuse after Flush.
|
||||
slot := b.Reserve(32)
|
||||
if cap(slot) != 32 {
|
||||
t.Fatalf("after Flush Reserve wrong cap: %d", cap(slot))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
|
||||
fw := &fakeBatchWriter{}
|
||||
b := NewSendBatch(fw, 3, 8)
|
||||
ap := netip.MustParseAddrPort("10.0.0.1:80")
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
s := b.Reserve(8)
|
||||
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
|
||||
b.Commit(pkt, ap)
|
||||
}
|
||||
if _, err := b.Flush(); err != nil {
|
||||
t.Fatalf("Flush: %v", err)
|
||||
}
|
||||
|
||||
for i, buf := range fw.bufs {
|
||||
if buf[0] != byte(0xA0+i) || buf[1] != byte(0xB0+i) {
|
||||
t.Errorf("slot %d corrupted: %x", i, buf)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendBatchGrowPreservesCommitted(t *testing.T) {
|
||||
fw := &fakeBatchWriter{}
|
||||
// Tiny initial backing forces a grow on the second Reserve.
|
||||
b := NewSendBatch(fw, 1, 4)
|
||||
ap := netip.MustParseAddrPort("10.0.0.1:80")
|
||||
|
||||
s1 := b.Reserve(4)
|
||||
pkt1 := append(s1[:0], 0x11, 0x22, 0x33, 0x44)
|
||||
b.Commit(pkt1, ap)
|
||||
|
||||
s2 := b.Reserve(8) // exceeds remaining cap, triggers grow
|
||||
pkt2 := append(s2[:0], 0xA, 0xB, 0xC, 0xD, 0xE)
|
||||
b.Commit(pkt2, ap)
|
||||
|
||||
// pkt1 must still be intact even though backing reallocated.
|
||||
if pkt1[0] != 0x11 || pkt1[3] != 0x44 {
|
||||
t.Fatalf("first packet corrupted by grow: %x", pkt1)
|
||||
}
|
||||
|
||||
if _, err := b.Flush(); err != nil {
|
||||
t.Fatalf("Flush: %v", err)
|
||||
}
|
||||
if len(fw.bufs) != 2 {
|
||||
t.Fatalf("got %d bufs want 2", len(fw.bufs))
|
||||
}
|
||||
if fw.bufs[0][0] != 0x11 || fw.bufs[0][3] != 0x44 {
|
||||
t.Errorf("first packet on the wire: %x", fw.bufs[0])
|
||||
}
|
||||
if fw.bufs[1][0] != 0xA || fw.bufs[1][4] != 0xE {
|
||||
t.Errorf("second packet on the wire: %x", fw.bufs[1])
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,345 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
)
|
||||
|
||||
// ipProtoUDP is the IANA protocol number for UDP.
|
||||
const ipProtoUDP = 17
|
||||
|
||||
// udpCoalesceBufSize caps total bytes per UDP superpacket. Mirrors the
|
||||
// kernel's gso_max_size; payloads beyond this are emitted as-is.
|
||||
const udpCoalesceBufSize = 65535
|
||||
|
||||
// udpCoalesceMaxSegs caps how many segments we'll coalesce. Kernel UDP-GSO
|
||||
// accepts up to 64 segments per skb (UDP_MAX_SEGMENTS); stay under that.
|
||||
const udpCoalesceMaxSegs = 64
|
||||
|
||||
// udpSlot is one entry in the UDPCoalescer's ordered event queue.
|
||||
type udpSlot struct {
|
||||
verbatim bool
|
||||
// rawPkt is borrowed: the whole packet for verbatim slots, the seed
|
||||
// packet for coalesce slots. A coalesce slot that never grows past one
|
||||
// segment is emitted from rawPkt so its original (already valid) L4
|
||||
// checksum ships DATA_VALID instead of making the kernel recompute it.
|
||||
// A multi-segment slot's superpacket header is rawPkt's, patched in place at flush.
|
||||
rawPkt []byte
|
||||
|
||||
fk flowKey
|
||||
hdrLen int
|
||||
ipHdrLen int
|
||||
isV6 bool
|
||||
gsoSize int // per-segment UDP payload length
|
||||
numSeg int
|
||||
totalPay int
|
||||
payIovs [][]byte
|
||||
}
|
||||
|
||||
// UDPCoalescer accumulates adjacent in-flow UDP datagrams across multiple
|
||||
// concurrent flows and emits each flow's run as a single GSO_UDP_L4 superpacket via tio.GSOWriter.
|
||||
// Preserves the in-flow order of packets as they are Commit-ed
|
||||
//
|
||||
// Owns no locks; one coalescer per TUN write queue.
|
||||
type UDPCoalescer struct {
|
||||
w tio.GSOWriter
|
||||
slots []*udpSlot
|
||||
openSlots map[flowKey]*udpSlot
|
||||
// lastSlot caches the most recently touched open slot; see the
|
||||
// TCPCoalescer field of the same name. Single-flow QUIC bulk is the
|
||||
// dominant USO workload, and multi-flow arrival comes in GRO runs, so
|
||||
// the fk compare beats the map's 38-byte key hash on most packets.
|
||||
// Kept in lockstep with openSlots: nil whenever the slot it pointed at
|
||||
// is removed.
|
||||
lastSlot *udpSlot
|
||||
pool []*udpSlot
|
||||
}
|
||||
|
||||
func NewUDPCoalescer(w io.Writer) *UDPCoalescer {
|
||||
gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
return &UDPCoalescer{
|
||||
w: gw,
|
||||
slots: make([]*udpSlot, 0, initialSlots),
|
||||
openSlots: make(map[flowKey]*udpSlot, initialSlots),
|
||||
pool: make([]*udpSlot, 0, initialSlots),
|
||||
}
|
||||
}
|
||||
|
||||
// parsedUDP holds the fields extracted from a single parse so later steps
|
||||
// (admission, slot lookup, canAppend) don't re-walk the header.
|
||||
type parsedUDP struct {
|
||||
fk flowKey
|
||||
ipHdrLen int
|
||||
hdrLen int // ipHdrLen + 8
|
||||
payLen int
|
||||
}
|
||||
|
||||
// parseAt extracts the flow key and IP/UDP offsets for a packet the dispatcher already knows is
|
||||
// UDP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on
|
||||
// entry and is filled in place. Returns false for malformed input or any shape that must not
|
||||
// coalesce (IPv4 options/fragmentation, IPv6 extension headers).
|
||||
func (p *parsedUDP) parseAt(pkt []byte, ipHdrLen int) bool {
|
||||
trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return p.parseTail(trimmed, ipHdrLen)
|
||||
}
|
||||
|
||||
// parseTail layers the UDP-header parse on a validated IP prologue. pkt is the trimmed packet;
|
||||
// fk's addresses are already filled.
|
||||
func (p *parsedUDP) parseTail(pkt []byte, ipHdrLen int) bool {
|
||||
if len(pkt) < ipHdrLen+8 {
|
||||
return false
|
||||
}
|
||||
// UDP `length` field: must equal IP-derived length-of-UDP-header-plus-payload.
|
||||
udpLen := int(binary.BigEndian.Uint16(pkt[ipHdrLen+4 : ipHdrLen+6]))
|
||||
if udpLen < 8 || udpLen > len(pkt)-ipHdrLen {
|
||||
return false
|
||||
}
|
||||
p.ipHdrLen = ipHdrLen
|
||||
p.hdrLen = ipHdrLen + 8
|
||||
p.payLen = udpLen - 8
|
||||
p.fk.sport = binary.BigEndian.Uint16(pkt[ipHdrLen : ipHdrLen+2])
|
||||
p.fk.dport = binary.BigEndian.Uint16(pkt[ipHdrLen+2 : ipHdrLen+4])
|
||||
return true
|
||||
}
|
||||
|
||||
// sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips
|
||||
// hashing the 38-byte key when no chains are open.
|
||||
func (c *UDPCoalescer) sealFlow(fk flowKey) {
|
||||
if len(c.openSlots) == 0 {
|
||||
return
|
||||
}
|
||||
if last := c.lastSlot; last != nil && last.fk == fk {
|
||||
c.lastSlot = nil
|
||||
}
|
||||
delete(c.openSlots, fk)
|
||||
}
|
||||
|
||||
// commitStaged commits one staged packet dispatch routed to this lane. A shape the lane cannot
|
||||
// coalesce (any fragmentation, unparseable header) seals every open chain — its flow is unknown —
|
||||
// and rides the lane as an in-lane verbatim, still in transmission order.
|
||||
func (c *UDPCoalescer) commitStaged(sp stagedPacket) error {
|
||||
if sp.fragAny {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(sp.pkt)
|
||||
return nil
|
||||
}
|
||||
var info parsedUDP
|
||||
if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
|
||||
c.sealAllOpen()
|
||||
c.addVerbatim(sp.pkt)
|
||||
return nil
|
||||
}
|
||||
return c.commitParsed(sp.pkt, &info)
|
||||
}
|
||||
|
||||
// commitParsed commits one parsed UDP packet. The caller (dispatch, via parseAt) supplies a
|
||||
// valid parse so the header is not re-walked here.
|
||||
func (c *UDPCoalescer) commitParsed(pkt []byte, info *parsedUDP) error {
|
||||
// A zero-length UDP datagram (length == 8) is legal and must reach the TUN, but cannot be
|
||||
// coalesced.
|
||||
if info.payLen == 0 {
|
||||
c.sealFlow(info.fk)
|
||||
c.addVerbatim(pkt)
|
||||
return nil
|
||||
}
|
||||
// Cached-slot fast path; see the TCPCoalescer equivalent.
|
||||
var open *udpSlot
|
||||
if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||
open = last
|
||||
} else {
|
||||
open = c.openSlots[info.fk]
|
||||
}
|
||||
if open != nil {
|
||||
if c.canAppend(open, pkt, info) {
|
||||
if c.appendPayload(open, pkt, info) {
|
||||
// Chain closed (short segment): stop extending it.
|
||||
c.sealFlow(info.fk)
|
||||
} else {
|
||||
c.lastSlot = open
|
||||
}
|
||||
return nil
|
||||
}
|
||||
// Can't extend: evict it from openSlots and fall through to seed a
|
||||
// fresh slot.
|
||||
c.sealFlow(info.fk)
|
||||
}
|
||||
c.seed(pkt, info)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) Flush() error {
|
||||
var first error
|
||||
for _, s := range c.slots {
|
||||
var err error
|
||||
if s.verbatim || s.numSeg == 1 {
|
||||
// A slot that never grew is byte-identical to the packet it was
|
||||
// seeded from; ship the original so its valid checksum rides the
|
||||
// DATA_VALID path instead of paying a kernel software csum.
|
||||
_, err = c.w.Write(s.rawPkt)
|
||||
} else {
|
||||
err = c.flushSlot(s)
|
||||
}
|
||||
if err != nil && first == nil {
|
||||
first = err
|
||||
}
|
||||
c.release(s)
|
||||
}
|
||||
clear(c.slots)
|
||||
c.slots = c.slots[:0]
|
||||
clear(c.openSlots)
|
||||
c.lastSlot = nil
|
||||
return first
|
||||
}
|
||||
|
||||
// sealAllOpen closes every open coalesce chain. Called for unparseable packets: the flow key is
|
||||
// unknown, so any open chain could otherwise absorb later data and emit it ahead of this packet.
|
||||
func (c *UDPCoalescer) sealAllOpen() {
|
||||
clear(c.openSlots)
|
||||
c.lastSlot = nil
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) addVerbatim(pkt []byte) {
|
||||
s := c.take()
|
||||
s.verbatim = true
|
||||
s.rawPkt = pkt
|
||||
c.slots = append(c.slots, s)
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) seed(pkt []byte, info *parsedUDP) {
|
||||
if info.hdrLen+info.payLen > udpCoalesceBufSize {
|
||||
// Pathological shape that can't ride a superpacket; emit as-is. No chain for this flow can
|
||||
// be open here (commitParsed evicts before seeding), so sealFlow is defense in depth
|
||||
// against a stale cache entry absorbing later data.
|
||||
c.sealFlow(info.fk)
|
||||
c.addVerbatim(pkt)
|
||||
return
|
||||
}
|
||||
s := c.take()
|
||||
s.verbatim = false
|
||||
// rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is
|
||||
// the superpacket header flushSlot patches in place.
|
||||
s.rawPkt = pkt
|
||||
s.hdrLen = info.hdrLen
|
||||
s.ipHdrLen = info.ipHdrLen
|
||||
s.isV6 = info.fk.isV6
|
||||
s.fk = info.fk
|
||||
s.gsoSize = info.payLen
|
||||
s.numSeg = 1
|
||||
s.totalPay = info.payLen
|
||||
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||
c.slots = append(c.slots, s)
|
||||
c.openSlots[info.fk] = s
|
||||
c.lastSlot = s
|
||||
}
|
||||
|
||||
// canAppend reports whether info's packet extends the slot's seed.
|
||||
// Kernel UDP-GSO requires every segment except possibly the last to be
|
||||
// exactly gsoSize, and the last may be shorter (≤ gsoSize).
|
||||
func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info *parsedUDP) bool {
|
||||
if info.hdrLen != s.hdrLen {
|
||||
return false
|
||||
}
|
||||
if s.numSeg >= udpCoalesceMaxSegs {
|
||||
return false
|
||||
}
|
||||
if info.payLen > s.gsoSize {
|
||||
return false
|
||||
}
|
||||
if s.hdrLen+s.totalPay+info.payLen > udpCoalesceBufSize {
|
||||
return false
|
||||
}
|
||||
// Header reads use rawPkt, which is never mutated before flush. A closed chain never reaches
|
||||
// here; closing removes the slot from openSlots, the only path in.
|
||||
if !s.isV6 && !ipv4CanCoalesceID(s.rawPkt, pkt, s.numSeg) {
|
||||
return false
|
||||
}
|
||||
if !udpHeadersMatch(s.rawPkt[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// appendPayload folds info's packet into s and reports whether the chain is now closed: kernel
|
||||
// UDP-GSO requires every segment but the last to be exactly gsoSize, so a short segment must be
|
||||
// the final one. The caller must deregister a closed slot from openSlots.
|
||||
func (c *UDPCoalescer) appendPayload(s *udpSlot, pkt []byte, info *parsedUDP) bool {
|
||||
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||
s.numSeg++
|
||||
s.totalPay += info.payLen
|
||||
return info.payLen < s.gsoSize
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) take() *udpSlot {
|
||||
if n := len(c.pool); n > 0 {
|
||||
s := c.pool[n-1]
|
||||
c.pool[n-1] = nil
|
||||
c.pool = c.pool[:n-1]
|
||||
return s
|
||||
}
|
||||
return &udpSlot{}
|
||||
}
|
||||
|
||||
func (c *UDPCoalescer) release(s *udpSlot) {
|
||||
// Reset every field, identity ones included; see TCPCoalescer.release.
|
||||
clear(s.payIovs)
|
||||
*s = udpSlot{payIovs: s.payIovs[:0]}
|
||||
c.pool = append(c.pool, s)
|
||||
}
|
||||
|
||||
// flushSlot patches the IP header total length / IPv6 payload length and
|
||||
// the UDP length to the *total* across all coalesced segments, then seeds
|
||||
// the UDP checksum field with the pseudo-header partial (single-fold, not
|
||||
// inverted) per virtio NEEDS_CSUM. The patches land in place in rawPkt; the
|
||||
// slot is released right after, so nothing re-reads the patched header.
|
||||
func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
|
||||
hdr := s.rawPkt[:s.hdrLen]
|
||||
total := s.hdrLen + s.totalPay // full IP+UDP+all_payloads bytes
|
||||
l4Len := total - s.ipHdrLen // total UDP (8 + sum of payloads)
|
||||
|
||||
if s.isV6 {
|
||||
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
|
||||
} else {
|
||||
binary.BigEndian.PutUint16(hdr[2:4], uint16(total))
|
||||
hdr[10] = 0
|
||||
hdr[11] = 0
|
||||
binary.BigEndian.PutUint16(hdr[10:12], ipv4HdrChecksum(hdr[:s.ipHdrLen]))
|
||||
}
|
||||
|
||||
// UDP length field (offset 4 inside the UDP header) = total UDP size.
|
||||
binary.BigEndian.PutUint16(hdr[s.ipHdrLen+4:s.ipHdrLen+6], uint16(l4Len))
|
||||
|
||||
var psum uint32
|
||||
if s.isV6 {
|
||||
psum = pseudoSumIPv6(hdr[8:24], hdr[24:40], ipProtoUDP, l4Len)
|
||||
} else {
|
||||
psum = pseudoSumIPv4(hdr[12:16], hdr[16:20], ipProtoUDP, l4Len)
|
||||
}
|
||||
udpCsumOff := s.ipHdrLen + 6
|
||||
binary.BigEndian.PutUint16(hdr[udpCsumOff:udpCsumOff+2], foldOnceNoInvert(psum))
|
||||
|
||||
return c.w.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoUDP)
|
||||
}
|
||||
|
||||
// udpHeadersMatch compares two IP+UDP header prefixes for byte-equality on
|
||||
// every field that must be identical across coalesced segments
|
||||
func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
if !ipHeadersMatch(a, b, isV6) {
|
||||
return false
|
||||
}
|
||||
// UDP: compare sport+dport ([0:4]). Skip length [4:6] and checksum [6:8]:
|
||||
// length varies (we rewrite at flush) and the checksum will be redone.
|
||||
udp := ipHdrLen
|
||||
return bytes.Equal(a[udp:udp+4], b[udp:udp+4])
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
// buildUDPv4BulkFlow returns n equal-size datagrams on one flow — the
|
||||
// steady state for single-flow QUIC bulk, the workload USO exists for.
|
||||
func buildUDPv4BulkFlow(n, payloadLen int) [][]byte {
|
||||
pay := make([]byte, payloadLen)
|
||||
pkts := make([][]byte, n)
|
||||
for i := range pkts {
|
||||
pkts[i] = buildUDPv4(40000, 443, pay)
|
||||
}
|
||||
return pkts
|
||||
}
|
||||
|
||||
// buildUDPv4RunInterleaved mirrors buildTCPv4RunInterleaved: nFlows*perFlow
|
||||
// datagrams arriving in GRO-burst runs of runLen per flow.
|
||||
func buildUDPv4RunInterleaved(nFlows, perFlow, runLen, payloadLen int) [][]byte {
|
||||
pay := make([]byte, payloadLen)
|
||||
pkts := make([][]byte, 0, nFlows*perFlow)
|
||||
for done := 0; done < perFlow; done += runLen {
|
||||
for f := range nFlows {
|
||||
sport := uint16(40000 + f)
|
||||
for range runLen {
|
||||
pkts = append(pkts, buildUDPv4(sport, 443, pay))
|
||||
}
|
||||
}
|
||||
}
|
||||
return pkts
|
||||
}
|
||||
|
||||
// runUDPCommitBench drives UDPCoalescer.Commit over pkts batchSize at a
|
||||
// time, flushing between batches, and reports per-packet cost.
|
||||
func runUDPCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||
b.Helper()
|
||||
c := newTestUDPCoalescer(b, nopTunWriter{})
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkts[0])))
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
pkt := pkts[i%len(pkts)]
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if (i+1)%batchSize == 0 {
|
||||
if err := c.Flush(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
_ = c.Flush()
|
||||
}
|
||||
|
||||
// BenchmarkUDPCommitSingleFlow is the single-flow bulk steady state.
|
||||
func BenchmarkUDPCommitSingleFlow(b *testing.B) {
|
||||
pkts := buildUDPv4BulkFlow(udpCoalesceMaxSegs, 1200)
|
||||
runUDPCommitBench(b, pkts, udpCoalesceMaxSegs)
|
||||
}
|
||||
|
||||
// BenchmarkUDPCommitInterleaved4 is the adversarial per-packet round-robin.
|
||||
func BenchmarkUDPCommitInterleaved4(b *testing.B) {
|
||||
pkts := buildUDPv4RunInterleaved(4, udpCoalesceMaxSegs, 1, 1200)
|
||||
runUDPCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
|
||||
// BenchmarkUDPCommitRunInterleaved4 is 4 flows in GRO-burst runs of 16.
|
||||
func BenchmarkUDPCommitRunInterleaved4(b *testing.B) {
|
||||
pkts := buildUDPv4RunInterleaved(4, udpCoalesceMaxSegs, 16, 1200)
|
||||
runUDPCommitBench(b, pkts, len(pkts))
|
||||
}
|
||||
@@ -0,0 +1,536 @@
|
||||
package batch
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// buildUDPv4 builds a minimal IPv4+UDP packet with the given payload and ports.
|
||||
func buildUDPv4(sport, dport uint16, payload []byte) []byte {
|
||||
const ipHdrLen = 20
|
||||
const udpHdrLen = 8
|
||||
total := ipHdrLen + udpHdrLen + len(payload)
|
||||
pkt := make([]byte, total)
|
||||
|
||||
pkt[0] = 0x45
|
||||
pkt[1] = 0x00
|
||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(total))
|
||||
binary.BigEndian.PutUint16(pkt[4:6], 0)
|
||||
binary.BigEndian.PutUint16(pkt[6:8], 0x4000)
|
||||
pkt[8] = 64
|
||||
pkt[9] = ipProtoUDP
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||
|
||||
binary.BigEndian.PutUint16(pkt[20:22], sport)
|
||||
binary.BigEndian.PutUint16(pkt[22:24], dport)
|
||||
binary.BigEndian.PutUint16(pkt[24:26], uint16(udpHdrLen+len(payload)))
|
||||
binary.BigEndian.PutUint16(pkt[26:28], 0)
|
||||
|
||||
copy(pkt[28:], payload)
|
||||
return pkt
|
||||
}
|
||||
|
||||
// buildUDPv6 builds a minimal IPv6+UDP packet.
|
||||
func buildUDPv6(sport, dport uint16, payload []byte) []byte {
|
||||
const ipHdrLen = 40
|
||||
const udpHdrLen = 8
|
||||
total := ipHdrLen + udpHdrLen + len(payload)
|
||||
pkt := make([]byte, total)
|
||||
|
||||
pkt[0] = 0x60
|
||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(udpHdrLen+len(payload)))
|
||||
pkt[6] = ipProtoUDP
|
||||
pkt[7] = 64
|
||||
pkt[8] = 0xfe
|
||||
pkt[9] = 0x80
|
||||
pkt[23] = 1
|
||||
pkt[24] = 0xfe
|
||||
pkt[25] = 0x80
|
||||
pkt[39] = 2
|
||||
|
||||
binary.BigEndian.PutUint16(pkt[40:42], sport)
|
||||
binary.BigEndian.PutUint16(pkt[42:44], dport)
|
||||
binary.BigEndian.PutUint16(pkt[44:46], uint16(udpHdrLen+len(payload)))
|
||||
binary.BigEndian.PutUint16(pkt[46:48], 0)
|
||||
|
||||
copy(pkt[48:], payload)
|
||||
return pkt
|
||||
}
|
||||
|
||||
// newTestUDPCoalescer builds a coalescer over w and fails the test if w can't
|
||||
// do USO. See newTestTCPCoalescer.
|
||||
func newTestUDPCoalescer(tb testing.TB, w io.Writer) *UDPCoalescer {
|
||||
tb.Helper()
|
||||
c := NewUDPCoalescer(w)
|
||||
if c == nil {
|
||||
tb.Fatal("NewUDPCoalescer: writer does not support USO")
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
// TestNewUDPCoalescerRefusesWhenGSOUnavailable mirrors the TCP precondition:
|
||||
// no USO, no coalescer.
|
||||
func TestNewUDPCoalescerRefusesWhenGSOUnavailable(t *testing.T) {
|
||||
if c := NewUDPCoalescer(&fakeTunWriter{gsoEnabled: false}); c != nil {
|
||||
t.Fatalf("want nil for a non-USO writer, got %v", c)
|
||||
}
|
||||
if c := NewUDPCoalescer(&plainOnlyWriter{}); c != nil {
|
||||
t.Fatalf("want nil for a plain writer, got %v", c)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
// ICMP packet
|
||||
pkt := make([]byte, 28)
|
||||
pkt[0] = 0x45
|
||||
binary.BigEndian.PutUint16(pkt[2:4], 28)
|
||||
pkt[9] = 1
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("ICMP must pass through unchanged: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
}
|
||||
|
||||
func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pkt := buildUDPv4(1000, 53, make([]byte, 800))
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// A slot that never grew past one datagram flushes as a plain Write of
|
||||
// the original packet bytes: the original (already valid) checksum
|
||||
// ships via the DATA_VALID path, so the kernel does no csum work.
|
||||
// WriteGSO is reserved for slots that actually coalesced (>=2 segs).
|
||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("single-seg flush: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
if !bytes.Equal(w.writes[0], pkt) {
|
||||
t.Errorf("plain write not byte-identical to committed packet: got %d bytes want %d", len(w.writes[0]), len(pkt))
|
||||
}
|
||||
}
|
||||
|
||||
func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
for i := 0; i < 3; i++ {
|
||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 1 {
|
||||
t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
g := w.gsoWrites[0]
|
||||
if g.gsoSize != 1200 {
|
||||
t.Errorf("gsoSize=%d want 1200", g.gsoSize)
|
||||
}
|
||||
if len(g.pays) != 3 {
|
||||
t.Errorf("pay count=%d want 3", len(g.pays))
|
||||
}
|
||||
if g.csumStart != 20 {
|
||||
t.Errorf("csumStart=%d want 20", g.csumStart)
|
||||
}
|
||||
// IP totalLen and UDP length must be the TOTAL across all segments —
|
||||
// the kernel's ip_rcv_core trims skbs to iph->tot_len, so a per-segment
|
||||
// value would silently drop everything but the first segment. Total =
|
||||
// IP(20) + UDP(8) + 3*1200 = 3628.
|
||||
gotTotalLen := binary.BigEndian.Uint16(g.hdr[2:4])
|
||||
if gotTotalLen != 3628 {
|
||||
t.Errorf("ipv4 total_len=%d want 3628 (must be total across segments)", gotTotalLen)
|
||||
}
|
||||
gotUDPLen := binary.BigEndian.Uint16(g.hdr[20+4 : 20+6])
|
||||
if gotUDPLen != 8+3*1200 {
|
||||
t.Errorf("udp len=%d want %d", gotUDPLen, 8+3*1200)
|
||||
}
|
||||
}
|
||||
|
||||
// Last segment may be shorter, sealing the chain.
|
||||
func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
full := make([]byte, 1200)
|
||||
tail := make([]byte, 600)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(1000, 53, tail)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// A 4th packet, even same-sized, must NOT join — chain is sealed.
|
||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// The sealed 3-datagram chain is a real superpacket; the re-seed stays
|
||||
// single-segment and flushes as a plain write of the original packet.
|
||||
if len(w.gsoWrites) != 1 || len(w.writes) != 1 {
|
||||
t.Fatalf("want 1 gso (sealed) + 1 plain (new seed), got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
if len(w.gsoWrites[0].pays) != 3 {
|
||||
t.Errorf("super: want 3 pays, got %d", len(w.gsoWrites[0].pays))
|
||||
}
|
||||
if got, want := len(w.writes[0]), 20+8+1200; got != want {
|
||||
t.Errorf("re-seed plain write len=%d want %d", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// A larger-than-gsoSize packet cannot extend the slot — it reseeds.
|
||||
func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 1200))); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Both seeds stay single-segment → two plain writes in arrival order.
|
||||
if len(w.writes) != 2 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("want 2 separate plain writes, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
for i, want := range []int{20 + 8 + 800, 20 + 8 + 1200} {
|
||||
if len(w.writes[i]) != want {
|
||||
t.Errorf("write %d len=%d want %d", i, len(w.writes[i]), want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Different 5-tuples must not coalesce.
|
||||
func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 800)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(2000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(2000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Two flows × 2 datagrams each = 2 superpackets of 2 segments.
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 gso writes (one per flow), got %d", len(w.gsoWrites))
|
||||
}
|
||||
for i, g := range w.gsoWrites {
|
||||
if len(g.pays) != 2 {
|
||||
t.Errorf("super %d: want 2 pays, got %d", i, len(g.pays))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Caps at udpCoalesceMaxSegs.
|
||||
func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 100)
|
||||
for i := 0; i < udpCoalesceMaxSegs+5; i++ {
|
||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// First superpacket holds udpCoalesceMaxSegs segments; the spillover
|
||||
// reseeds a new one.
|
||||
if len(w.gsoWrites) != 2 {
|
||||
t.Fatalf("want 2 gso writes (cap then reseed), got %d", len(w.gsoWrites))
|
||||
}
|
||||
if len(w.gsoWrites[0].pays) != udpCoalesceMaxSegs {
|
||||
t.Errorf("first super: pays=%d want %d", len(w.gsoWrites[0].pays), udpCoalesceMaxSegs)
|
||||
}
|
||||
if len(w.gsoWrites[1].pays) != 5 {
|
||||
t.Errorf("second super: pays=%d want 5", len(w.gsoWrites[1].pays))
|
||||
}
|
||||
}
|
||||
|
||||
// Differing IP ECN codepoints must not coalesce: udpHeadersMatch compares
|
||||
// the full ToS byte (matching kernel GRO). A CE-marked datagram mid-run
|
||||
// seals the Not-ECT chain and reseeds; the trailing Not-ECT datagram
|
||||
// reseeds again. All three stay single-segment, so each ships as a plain
|
||||
// write of its original bytes, keeping its own codepoint.
|
||||
func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 800)
|
||||
pkt0 := buildUDPv4(1000, 53, pay) // ECN=00 (Not-ECT)
|
||||
pkt1 := buildUDPv4(1000, 53, pay)
|
||||
pkt1[1] = 0x03 // CE
|
||||
pkt2 := buildUDPv4(1000, 53, pay) // ECN=00 again
|
||||
for _, p := range [][]byte{pkt0, pkt1, pkt2} {
|
||||
if err := c.Commit(p); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 3 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("want 3 separate plain writes (differing ECN), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
wantECN := []byte{0x00, 0x03, 0x00}
|
||||
for i, p := range w.writes {
|
||||
if got := p[1] & 0x03; got != wantECN[i] {
|
||||
t.Errorf("write %d ECN=%#x want %#x", i, got, wantECN[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// IPv6 path: same flow, equal-sized → coalesced.
|
||||
func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
for i := 0; i < 3; i++ {
|
||||
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 1 {
|
||||
t.Fatalf("want 1 gso write, got %d", len(w.gsoWrites))
|
||||
}
|
||||
g := w.gsoWrites[0]
|
||||
if !g.isV6 {
|
||||
t.Errorf("expected v6 write")
|
||||
}
|
||||
if g.csumStart != 40 {
|
||||
t.Errorf("csumStart=%d want 40", g.csumStart)
|
||||
}
|
||||
// IPv6 payload_len and UDP length must be TOTAL — kernel's
|
||||
// ip6_rcv_core trims to payload_len + ipv6 hdr size. Total UDP = 8 +
|
||||
// 3*1200 = 3608.
|
||||
gotPlen := binary.BigEndian.Uint16(g.hdr[4:6])
|
||||
if gotPlen != 8+3*1200 {
|
||||
t.Errorf("ipv6 payload_len=%d want %d (must be total)", gotPlen, 8+3*1200)
|
||||
}
|
||||
gotUDPLen := binary.BigEndian.Uint16(g.hdr[40+4 : 40+6])
|
||||
if gotUDPLen != 8+3*1200 {
|
||||
t.Errorf("udp len=%d want %d", gotUDPLen, 8+3*1200)
|
||||
}
|
||||
}
|
||||
|
||||
// DSCP differences must reseed: udpHeadersMatch compares the full ToS byte.
|
||||
func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 800)
|
||||
pkt0 := buildUDPv4(1000, 53, pay)
|
||||
pkt1 := buildUDPv4(1000, 53, pay)
|
||||
pkt1[1] = 0xb8 // EF DSCP, ECN=0
|
||||
if err := c.Commit(pkt0); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(pkt1); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Both seeds stay single-segment → two plain writes, no gso.
|
||||
if len(w.writes) != 2 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("want 2 separate plain writes (different DSCP), got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
}
|
||||
|
||||
// Fragmented IPv4 must not be coalesced.
|
||||
func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
||||
binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("frag must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
}
|
||||
|
||||
// A zero-length UDP datagram (UDP length == 8, no payload) is legal and
|
||||
// must be delivered as a plain single datagram — never coalesced. Seeding
|
||||
// it into a GSO slot stores an empty payload iovec that panics WriteGSO
|
||||
// (index-out-of-range on &pay[0]); this is a remote DoS if we ever let it
|
||||
// reach the GSO path. Regression: must not panic and must be written.
|
||||
func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pkt := buildUDPv4(1000, 53, nil) // UDP length 8, zero payload
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("zero-length UDP must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes[0]) != len(pkt) {
|
||||
t.Errorf("delivered %d bytes, want the whole %d-byte datagram", len(w.writes[0]), len(pkt))
|
||||
}
|
||||
}
|
||||
|
||||
// IPv6 zero-length UDP datagram: same verbatim contract as v4.
|
||||
func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pkt := buildUDPv6(1000, 53, nil) // UDP length 8, zero payload
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("zero-length IPv6 UDP must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
if len(w.writes[0]) != len(pkt) {
|
||||
t.Errorf("delivered %d bytes, want the whole %d-byte datagram", len(w.writes[0]), len(pkt))
|
||||
}
|
||||
}
|
||||
|
||||
// A zero-length datagram arriving mid-flow must seal the open chain so the
|
||||
// datagram after it seeds a fresh superpacket *after* the empty one on the
|
||||
// wire — per-flow arrival order (full, empty, full) must be preserved.
|
||||
func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
full := make([]byte, 800)
|
||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(1000, 53, nil)); err != nil { // zero-length
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// The empty datagram sealed the first slot, so the trailing full packet
|
||||
// can't join it. All three emit as plain writes (the two full datagrams
|
||||
// stayed single-segment; the empty one is verbatim) in per-flow
|
||||
// arrival order: full, empty, full.
|
||||
if len(w.writes) != 3 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("want 3 plain writes, got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
|
||||
}
|
||||
for i, want := range []int{20 + 8 + 800, 20 + 8, 20 + 8 + 800} {
|
||||
if len(w.writes[i]) != want {
|
||||
t.Errorf("write %d len=%d want %d (order full, empty, full)", i, len(w.writes[i]), want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// IPv4 with options is not admissible (we require IHL=5).
|
||||
func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
||||
pkt[0] = 0x46 // IHL = 6 (24-byte IPv4 header — has options)
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("ipv4-with-options must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
}
|
||||
|
||||
// TestUDPCoalescerNonAtomicSequentialIDsCoalesce mirrors the TCP rule: DF
|
||||
// clear is fine as long as the IDs already run seed+1 per datagram, so
|
||||
// kernel USO's re-stamp reproduces them.
|
||||
func TestUDPCoalescerNonAtomicSequentialIDsCoalesce(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
|
||||
for i := range 2 {
|
||||
pkt := buildUDPv4(40000, 443, pay)
|
||||
setIPv4ID(pkt, uint16(40+i), false)
|
||||
if err := c.Commit(pkt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.gsoWrites) != 1 || len(w.gsoWrites[0].pays) != 2 {
|
||||
t.Fatalf("sequential-ID DF=0 datagrams must coalesce: gso=%d", len(w.gsoWrites))
|
||||
}
|
||||
}
|
||||
|
||||
// TestUDPCoalescerNonAtomicIDGapReseeds: an ID jump on a DF=0 flow breaks
|
||||
// the chain; each datagram stays a single-segment slot and flushes as a
|
||||
// plain write that keeps its own (meaningful) ID.
|
||||
func TestUDPCoalescerNonAtomicIDGapReseeds(t *testing.T) {
|
||||
w := &fakeTunWriter{gsoEnabled: true}
|
||||
c := newTestUDPCoalescer(t, w)
|
||||
pay := make([]byte, 1200)
|
||||
|
||||
p1 := buildUDPv4(40000, 443, pay)
|
||||
setIPv4ID(p1, 40, false)
|
||||
p2 := buildUDPv4(40000, 443, pay)
|
||||
setIPv4ID(p2, 50, false)
|
||||
|
||||
if err := c.Commit(p1); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Commit(p2); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(w.writes) != 2 || len(w.gsoWrites) != 0 {
|
||||
t.Fatalf("ID gap on DF=0 must reseed: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||
}
|
||||
for i, want := range []uint16{40, 50} {
|
||||
if id := binary.BigEndian.Uint16(w.writes[i][4:6]); id != want {
|
||||
t.Errorf("write %d: ID=%d want %d (must be preserved)", i, id, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
package checksum
|
||||
|
||||
import (
|
||||
"golang.org/x/sys/cpu"
|
||||
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||
)
|
||||
|
||||
//go:noescape
|
||||
func checksumAVX2(buf []byte, initial uint16) uint16
|
||||
|
||||
var hasAVX2 = cpu.X86.HasAVX2
|
||||
|
||||
// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with
|
||||
// initial. It is a drop-in replacement for gvisor's checksum.Checksum that
|
||||
// dispatches to a hand-written AVX2 routine on amd64 CPUs that support it,
|
||||
// falling back to gvisor's pure-Go implementation otherwise. The result
|
||||
// matches gvisor's bit-for-bit for any buffer length and initial seed.
|
||||
func Checksum(buf []byte, initial uint16) uint16 {
|
||||
if hasAVX2 {
|
||||
return checksumAVX2(buf, initial)
|
||||
}
|
||||
return gvisorchecksum.Checksum(buf, initial)
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
#include "textflag.h"
|
||||
|
||||
// func checksumAVX2(buf []byte, initial uint16) uint16
|
||||
//
|
||||
// Computes the RFC 1071 ones-complement sum of buf, seeded with initial.
|
||||
//
|
||||
// Algorithm: sum the buffer treating it as a stream of uint32s in machine
|
||||
// (little-endian) byte order, accumulating into 64-bit lanes (top 32 bits
|
||||
// hold cross-add carries — at 1 byte / lane / iter we have 32 bits of
|
||||
// headroom which is far more than the 16 KB/64 KB max practical inputs).
|
||||
// At the end we fold to 16 bits and byte-swap once to recover the on-wire
|
||||
// (big-endian) result. RFC 1071 §1.2.B byte-order independence makes this
|
||||
// equivalent to summing as 16-bit big-endian words.
|
||||
//
|
||||
// The ymm accumulators (Y4..Y7) hold 4 uint64 lanes each = 16 parallel
|
||||
// partial sums. The main loop loads 64 bytes per iter as four 16-byte
|
||||
// chunks, zero-extending each chunk's four uint32s into a ymm via
|
||||
// VPMOVZXDQ-from-memory, then VPADDQ into a separate accumulator per
|
||||
// chunk to break the dep chain. After the vector loop the lane sums are
|
||||
// horizontally reduced and merged with a scalar accumulator that handles
|
||||
// the trailing 0..63 bytes plus the (byte-swapped) initial seed.
|
||||
TEXT ·checksumAVX2(SB), NOSPLIT, $0-34
|
||||
MOVQ buf_base+0(FP), SI
|
||||
MOVQ buf_len+8(FP), CX
|
||||
MOVWQZX initial+24(FP), AX
|
||||
|
||||
// Pre-byteswap initial into the LE-summing space so it merges directly
|
||||
// with the rest of the accumulator. The final fold's bswap16 will undo
|
||||
// this and convert the whole result back to BE.
|
||||
XCHGB AH, AL
|
||||
|
||||
CMPQ CX, $32
|
||||
JLT scalar_tail
|
||||
|
||||
VPXOR Y4, Y4, Y4
|
||||
VPXOR Y5, Y5, Y5
|
||||
VPXOR Y6, Y6, Y6
|
||||
VPXOR Y7, Y7, Y7
|
||||
|
||||
CMPQ CX, $64
|
||||
JLT loop32
|
||||
|
||||
loop64:
|
||||
VPMOVZXDQ (SI), Y0
|
||||
VPMOVZXDQ 16(SI), Y1
|
||||
VPMOVZXDQ 32(SI), Y2
|
||||
VPMOVZXDQ 48(SI), Y3
|
||||
VPADDQ Y0, Y4, Y4
|
||||
VPADDQ Y1, Y5, Y5
|
||||
VPADDQ Y2, Y6, Y6
|
||||
VPADDQ Y3, Y7, Y7
|
||||
ADDQ $64, SI
|
||||
SUBQ $64, CX
|
||||
CMPQ CX, $64
|
||||
JGE loop64
|
||||
|
||||
loop32:
|
||||
CMPQ CX, $32
|
||||
JLT reduce_vec
|
||||
VPMOVZXDQ (SI), Y0
|
||||
VPMOVZXDQ 16(SI), Y1
|
||||
VPADDQ Y0, Y4, Y4
|
||||
VPADDQ Y1, Y5, Y5
|
||||
ADDQ $32, SI
|
||||
SUBQ $32, CX
|
||||
JMP loop32
|
||||
|
||||
reduce_vec:
|
||||
// Combine the four ymm accumulators into Y4.
|
||||
VPADDQ Y5, Y4, Y4
|
||||
VPADDQ Y7, Y6, Y6
|
||||
VPADDQ Y6, Y4, Y4
|
||||
|
||||
// Horizontally reduce Y4's four uint64 lanes to a single scalar.
|
||||
VEXTRACTI128 $1, Y4, X5
|
||||
VPADDQ X5, X4, X4
|
||||
VPSHUFD $0x4e, X4, X5
|
||||
VPADDQ X5, X4, X4
|
||||
VMOVQ X4, R8
|
||||
VZEROUPPER
|
||||
|
||||
ADDQ R8, AX
|
||||
ADCQ $0, AX
|
||||
|
||||
scalar_tail:
|
||||
// Handle remaining 0..63 bytes (or the entire buffer if it was < 32).
|
||||
CMPQ CX, $8
|
||||
JLT tail4
|
||||
|
||||
loop8:
|
||||
ADDQ (SI), AX
|
||||
ADCQ $0, AX
|
||||
ADDQ $8, SI
|
||||
SUBQ $8, CX
|
||||
CMPQ CX, $8
|
||||
JGE loop8
|
||||
|
||||
tail4:
|
||||
CMPQ CX, $4
|
||||
JLT tail2
|
||||
MOVL (SI), R8
|
||||
ADDQ R8, AX
|
||||
ADCQ $0, AX
|
||||
ADDQ $4, SI
|
||||
SUBQ $4, CX
|
||||
|
||||
tail2:
|
||||
CMPQ CX, $2
|
||||
JLT tail1
|
||||
MOVWQZX (SI), R8
|
||||
ADDQ R8, AX
|
||||
ADCQ $0, AX
|
||||
ADDQ $2, SI
|
||||
SUBQ $2, CX
|
||||
|
||||
tail1:
|
||||
TESTQ CX, CX
|
||||
JZ fold
|
||||
MOVBQZX (SI), R8
|
||||
ADDQ R8, AX
|
||||
ADCQ $0, AX
|
||||
|
||||
fold:
|
||||
// Fold the 64-bit accumulator to 16 bits via four rounds, mirroring
|
||||
// gvisor's reduce(). Each pair (split, add) halves the live width;
|
||||
// the truncation steps absorb the single bit that may be left over
|
||||
// after each add so the next round's bound holds.
|
||||
|
||||
// 64 → 33 bits.
|
||||
MOVQ AX, R8
|
||||
SHRQ $32, R8
|
||||
MOVL AX, AX
|
||||
ADDQ R8, AX
|
||||
|
||||
// 33 → 32 bits. AX += (AX>>32); truncate to 32. AX is now ≤ 0xFFFF_FFFF.
|
||||
MOVQ AX, R8
|
||||
SHRQ $32, R8
|
||||
ADDQ R8, AX
|
||||
MOVL AX, AX
|
||||
|
||||
// 32 → 17 bits.
|
||||
MOVQ AX, R8
|
||||
SHRQ $16, R8
|
||||
MOVWQZX AX, AX
|
||||
ADDQ R8, AX
|
||||
|
||||
// 17 → 16 bits. AX += (AX>>16); the trailing MOVW truncates bit 16.
|
||||
MOVQ AX, R8
|
||||
SHRQ $16, R8
|
||||
ADDQ R8, AX
|
||||
|
||||
// AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip
|
||||
// to big-endian to match the gvisor API contract.
|
||||
XCHGB AH, AL
|
||||
|
||||
MOVW AX, ret+32(FP)
|
||||
RET
|
||||
@@ -0,0 +1,12 @@
|
||||
package checksum
|
||||
|
||||
//go:noescape
|
||||
func checksumNEON(buf []byte, initial uint16) uint16
|
||||
|
||||
// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with
|
||||
// initial. It is a drop-in replacement for gvisor's checksum.Checksum
|
||||
// that dispatches to a hand-written NEON routine. NEON is mandatory in
|
||||
// armv8 so no feature check is needed.
|
||||
func Checksum(buf []byte, initial uint16) uint16 {
|
||||
return checksumNEON(buf, initial)
|
||||
}
|
||||
@@ -0,0 +1,143 @@
|
||||
#include "textflag.h"
|
||||
|
||||
// func checksumNEON(buf []byte, initial uint16) uint16
|
||||
//
|
||||
// Mirrors the algorithm in checksum_amd64.s: sum the buffer treating it as
|
||||
// a stream of uint32s in machine (little-endian) byte order, accumulating
|
||||
// into 64-bit lanes that have ample carry headroom; fold and byte-swap once
|
||||
// at the very end to recover the on-wire (big-endian) result.
|
||||
//
|
||||
// Each loop iteration loads 64 bytes via VLD1.P into V0..V3 (4 Q regs).
|
||||
// VUADDW takes the low two uint32 lanes of a Q reg, zero-extends them to
|
||||
// uint64, and adds them into a 2×uint64 accumulator; VUADDW2 does the same
|
||||
// for the high two lanes. Four ymm-equivalent accumulators (V8..V11) get
|
||||
// updated twice per iter to break the dep chain. Tail bytes go through a
|
||||
// scalar ADCS chain seeded with the byte-swapped initial.
|
||||
TEXT ·checksumNEON(SB), NOSPLIT, $0-34
|
||||
MOVD buf_base+0(FP), R0
|
||||
MOVD buf_len+8(FP), R1
|
||||
MOVHU initial+24(FP), R2
|
||||
|
||||
// Pre-byteswap initial into the LE-summing space so it merges directly
|
||||
// with the rest of the accumulator.
|
||||
REV16W R2, R2
|
||||
|
||||
MOVD ZR, R3 // scalar accumulator
|
||||
|
||||
CMP $32, R1
|
||||
BLT scalar_tail
|
||||
|
||||
VEOR V8.B16, V8.B16, V8.B16
|
||||
VEOR V9.B16, V9.B16, V9.B16
|
||||
VEOR V10.B16, V10.B16, V10.B16
|
||||
VEOR V11.B16, V11.B16, V11.B16
|
||||
|
||||
CMP $64, R1
|
||||
BLT loop16_init
|
||||
|
||||
loop64:
|
||||
VLD1.P 64(R0), [V0.B16, V1.B16, V2.B16, V3.B16]
|
||||
VUADDW V0.S2, V8.D2, V8.D2
|
||||
VUADDW2 V0.S4, V9.D2, V9.D2
|
||||
VUADDW V1.S2, V10.D2, V10.D2
|
||||
VUADDW2 V1.S4, V11.D2, V11.D2
|
||||
VUADDW V2.S2, V8.D2, V8.D2
|
||||
VUADDW2 V2.S4, V9.D2, V9.D2
|
||||
VUADDW V3.S2, V10.D2, V10.D2
|
||||
VUADDW2 V3.S4, V11.D2, V11.D2
|
||||
SUB $64, R1, R1
|
||||
CMP $64, R1
|
||||
BGE loop64
|
||||
|
||||
loop16_init:
|
||||
CMP $16, R1
|
||||
BLT reduce_vec
|
||||
|
||||
loop16:
|
||||
VLD1.P 16(R0), [V0.B16]
|
||||
VUADDW V0.S2, V8.D2, V8.D2
|
||||
VUADDW2 V0.S4, V9.D2, V9.D2
|
||||
SUB $16, R1, R1
|
||||
CMP $16, R1
|
||||
BGE loop16
|
||||
|
||||
reduce_vec:
|
||||
// Combine the four accumulators into V8.
|
||||
VADD V9.D2, V8.D2, V8.D2
|
||||
VADD V11.D2, V10.D2, V10.D2
|
||||
VADD V10.D2, V8.D2, V8.D2
|
||||
|
||||
// Horizontal-add the two lanes of V8.D2 into a single uint64.
|
||||
VADDP V8.D2, V8.D2, V8.D2
|
||||
VMOV V8.D[0], R8
|
||||
|
||||
ADDS R8, R3, R3
|
||||
ADC ZR, R3, R3
|
||||
|
||||
scalar_tail:
|
||||
CMP $8, R1
|
||||
BLT tail4
|
||||
|
||||
loop8:
|
||||
MOVD.P 8(R0), R8
|
||||
ADDS R8, R3, R3
|
||||
ADC ZR, R3, R3
|
||||
SUB $8, R1, R1
|
||||
CMP $8, R1
|
||||
BGE loop8
|
||||
|
||||
tail4:
|
||||
CMP $4, R1
|
||||
BLT tail2
|
||||
MOVWU.P 4(R0), R8
|
||||
ADDS R8, R3, R3
|
||||
ADC ZR, R3, R3
|
||||
SUB $4, R1, R1
|
||||
|
||||
tail2:
|
||||
CMP $2, R1
|
||||
BLT tail1
|
||||
MOVHU.P 2(R0), R8
|
||||
ADDS R8, R3, R3
|
||||
ADC ZR, R3, R3
|
||||
SUB $2, R1, R1
|
||||
|
||||
tail1:
|
||||
CBZ R1, fold
|
||||
MOVBU (R0), R8
|
||||
ADDS R8, R3, R3
|
||||
ADC ZR, R3, R3
|
||||
|
||||
fold:
|
||||
// Merge the byte-swapped initial into our LE-form accumulator.
|
||||
ADDS R2, R3, R3
|
||||
ADC ZR, R3, R3
|
||||
|
||||
// 64 → 33 bits.
|
||||
LSR $32, R3, R8
|
||||
AND $0xffffffff, R3, R3
|
||||
ADD R8, R3, R3
|
||||
|
||||
// 33 → 32 (truncate after adding bit 32 back).
|
||||
LSR $32, R3, R8
|
||||
ADD R8, R3, R3
|
||||
AND $0xffffffff, R3, R3
|
||||
|
||||
// 32 → 17.
|
||||
LSR $16, R3, R8
|
||||
AND $0xffff, R3, R3
|
||||
ADD R8, R3, R3
|
||||
|
||||
// 17 → 16 (truncation absorbs bit 16 below).
|
||||
LSR $16, R3, R8
|
||||
ADD R8, R3, R3
|
||||
|
||||
// AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip
|
||||
// to big-endian to match the gvisor API contract. REV16W swaps bytes
|
||||
// within each 16-bit halfword of the low 32 bits, so it acts as a
|
||||
// 16-bit byte-swap on the live low 16.
|
||||
REV16W R3, R3
|
||||
AND $0xffff, R3, R3
|
||||
|
||||
MOVH R3, ret+32(FP)
|
||||
RET
|
||||
@@ -0,0 +1,10 @@
|
||||
//go:build !amd64 && !arm64
|
||||
|
||||
package checksum
|
||||
|
||||
import gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||
|
||||
// Checksum delegates to gvisor on architectures without a hand-written body.
|
||||
func Checksum(buf []byte, initial uint16) uint16 {
|
||||
return gvisorchecksum.Checksum(buf, initial)
|
||||
}
|
||||
@@ -0,0 +1,232 @@
|
||||
package checksum
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math/rand/v2"
|
||||
"testing"
|
||||
|
||||
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||
)
|
||||
|
||||
// archImpl names one checksum function under test. The per-arch
|
||||
// export_*_test.go files enumerate the hand-written implementations so the
|
||||
// suite compares each one against gvisor directly, regardless of which one
|
||||
// the public Checksum dispatches to on the running CPU. Testing only the
|
||||
// dispatcher was tautological wherever it resolved to the gvisor fallback
|
||||
// (non-AVX2 amd64, fallback architectures) — gvisor compared with itself,
|
||||
// assembly untested, suite green.
|
||||
type archImpl struct {
|
||||
name string
|
||||
fn func([]byte, uint16) uint16
|
||||
available bool
|
||||
}
|
||||
|
||||
// implsUnderTest is the public dispatcher plus every arch implementation.
|
||||
func implsUnderTest() []archImpl {
|
||||
return append([]archImpl{{name: "dispatch", fn: Checksum, available: true}}, archImpls...)
|
||||
}
|
||||
|
||||
// requireAvailable skips loudly when the running CPU can't execute an
|
||||
// implementation — visible in test output, unlike the old silent tautology.
|
||||
func requireAvailable(t *testing.T, impl archImpl) {
|
||||
t.Helper()
|
||||
if !impl.available {
|
||||
t.Skipf("%s not supported on this CPU; its assembly is NOT tested in this run", impl.name)
|
||||
}
|
||||
}
|
||||
|
||||
// TestChecksumMatchesGvisor walks lengths from 0 to 4096, with several initial
|
||||
// seeds and a handful of starting alignments, asserting that each local
|
||||
// implementation matches gvisor's reference bit-for-bit.
|
||||
func TestChecksumMatchesGvisor(t *testing.T) {
|
||||
for _, impl := range implsUnderTest() {
|
||||
t.Run(impl.name, func(t *testing.T) {
|
||||
requireAvailable(t, impl)
|
||||
rng := rand.New(rand.NewPCG(1, 2))
|
||||
const padFront = 16
|
||||
|
||||
// Random pool large enough for the longest case + alignment slop.
|
||||
pool := make([]byte, 4096+padFront)
|
||||
for i := range pool {
|
||||
pool[i] = byte(rng.Uint32())
|
||||
}
|
||||
|
||||
seeds := []uint16{0, 0x0001, 0xabcd, 0xffff, 0x1234, 0xfedc}
|
||||
offsets := []int{0, 1, 2, 3, 4, 5, 7, 8, 15, 16}
|
||||
|
||||
for length := 0; length <= 4096; length++ {
|
||||
for _, seed := range seeds {
|
||||
for _, off := range offsets {
|
||||
if off+length > len(pool) {
|
||||
continue
|
||||
}
|
||||
buf := pool[off : off+length]
|
||||
want := gvisorchecksum.Checksum(buf, seed)
|
||||
got := impl.fn(buf, seed)
|
||||
if got != want {
|
||||
t.Fatalf("len=%d off=%d seed=%#x: got %#04x want %#04x",
|
||||
length, off, seed, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestChecksumPatternedBuffers exercises specific byte patterns that have
|
||||
// historically tripped up checksum implementations: all-zero, all-0xff,
|
||||
// alternating, and ascending sequences.
|
||||
func TestChecksumPatternedBuffers(t *testing.T) {
|
||||
for _, impl := range implsUnderTest() {
|
||||
t.Run(impl.name, func(t *testing.T) {
|
||||
requireAvailable(t, impl)
|
||||
for length := 0; length <= 256; length++ {
|
||||
patterns := map[string][]byte{
|
||||
"zeros": make([]byte, length),
|
||||
"ones": bytes(length, 0xff),
|
||||
"alternating": pattern(length, []byte{0xa5, 0x5a}),
|
||||
"ascending": ascending(length),
|
||||
}
|
||||
for name, buf := range patterns {
|
||||
for _, seed := range []uint16{0, 0xffff, 0x8000} {
|
||||
want := gvisorchecksum.Checksum(buf, seed)
|
||||
got := impl.fn(buf, seed)
|
||||
if got != want {
|
||||
t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x",
|
||||
name, length, seed, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func bytes(n int, v byte) []byte {
|
||||
b := make([]byte, n)
|
||||
for i := range b {
|
||||
b[i] = v
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func pattern(n int, p []byte) []byte {
|
||||
b := make([]byte, n)
|
||||
for i := range b {
|
||||
b[i] = p[i%len(p)]
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func ascending(n int) []byte {
|
||||
b := make([]byte, n)
|
||||
for i := range b {
|
||||
b[i] = byte(i)
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
// TestChecksumTailPaths targets every combination of (SIMD body iterations,
|
||||
// trailing tail bytes) the asm handlers walk through. The tail handlers
|
||||
// peel off 8 → 4 → 2 → 1 byte chunks in turn; this test exercises each by
|
||||
// constructing lengths of the form 64*k + tail for tail ∈ [0, 63] and a
|
||||
// representative spread of k values, including k=0 (no main loop, all tail)
|
||||
// and k=1 (one main loop iter, then tail). It's explicit coverage for
|
||||
// payload sizes that are odd, not divisible by 4, by 8, or by 32.
|
||||
func TestChecksumTailPaths(t *testing.T) {
|
||||
for _, impl := range implsUnderTest() {
|
||||
t.Run(impl.name, func(t *testing.T) {
|
||||
requireAvailable(t, impl)
|
||||
rng := rand.New(rand.NewPCG(42, 17))
|
||||
const padFront = 16
|
||||
const maxK = 8
|
||||
|
||||
pool := make([]byte, 64*maxK+padFront+64)
|
||||
for i := range pool {
|
||||
pool[i] = byte(rng.Uint32())
|
||||
}
|
||||
|
||||
seeds := []uint16{0, 0xffff, 0xabcd}
|
||||
offsets := []int{0, 1, 3, 7, 15} // mix of aligned and odd starts
|
||||
|
||||
for k := 0; k <= maxK; k++ {
|
||||
for tail := 0; tail < 64; tail++ {
|
||||
length := 64*k + tail
|
||||
for _, seed := range seeds {
|
||||
for _, off := range offsets {
|
||||
if off+length > len(pool) {
|
||||
continue
|
||||
}
|
||||
buf := pool[off : off+length]
|
||||
want := gvisorchecksum.Checksum(buf, seed)
|
||||
got := impl.fn(buf, seed)
|
||||
if got != want {
|
||||
t.Fatalf("k=%d tail=%d (len=%d) off=%d seed=%#x: got %#04x want %#04x",
|
||||
k, tail, length, off, seed, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkChecksumTailSizes covers payload sizes that aren't clean multiples
|
||||
// of the SIMD body's 32-byte (amd64) or 16-byte (arm64) chunks, so the tail
|
||||
// handler is meaningfully on the hot path. Sizes are picked to either exercise
|
||||
// every tail branch (tiny lengths) or sit slightly off realistic packet
|
||||
// boundaries (e.g. 1499 = MTU − 1).
|
||||
func BenchmarkChecksumTailSizes(b *testing.B) {
|
||||
sizes := []int{
|
||||
1, 3, 7, 15, 31, // sub-SIMD; entire work is scalar tail
|
||||
33, 35, 47, 63, // one loop32 + assorted tails
|
||||
65, 95, 127, // one loop64 + assorted tails
|
||||
1447, 1471, 1499, 1501, // around MTU
|
||||
8191, 8193, // around USO
|
||||
65531, 65533, // near the kernel max
|
||||
}
|
||||
for _, size := range sizes {
|
||||
buf := make([]byte, size)
|
||||
for i := range buf {
|
||||
buf[i] = byte(i)
|
||||
}
|
||||
b.Run(fmt.Sprintf("size=%d/local", size), func(b *testing.B) {
|
||||
b.SetBytes(int64(size))
|
||||
for i := 0; i < b.N; i++ {
|
||||
_ = Checksum(buf, 0)
|
||||
}
|
||||
})
|
||||
b.Run(fmt.Sprintf("size=%d/gvisor", size), func(b *testing.B) {
|
||||
b.SetBytes(int64(size))
|
||||
for i := 0; i < b.N; i++ {
|
||||
_ = gvisorchecksum.Checksum(buf, 0)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkChecksum compares the local Checksum to gvisor's at sizes that
|
||||
// match real traffic: a TCP/IP header (60), a typical MSS (1448), a typical
|
||||
// USO size (8192), and the kernel's max GSO superpacket (65535).
|
||||
func BenchmarkChecksum(b *testing.B) {
|
||||
for _, size := range []int{60, 1448, 8192, 65535} {
|
||||
buf := make([]byte, size)
|
||||
for i := range buf {
|
||||
buf[i] = byte(i)
|
||||
}
|
||||
b.Run(fmt.Sprintf("size=%d/local", size), func(b *testing.B) {
|
||||
b.SetBytes(int64(size))
|
||||
for i := 0; i < b.N; i++ {
|
||||
_ = Checksum(buf, 0)
|
||||
}
|
||||
})
|
||||
b.Run(fmt.Sprintf("size=%d/gvisor", size), func(b *testing.B) {
|
||||
b.SetBytes(int64(size))
|
||||
for i := 0; i < b.N; i++ {
|
||||
_ = gvisorchecksum.Checksum(buf, 0)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
package checksum
|
||||
|
||||
// archImpls exposes every hand-written implementation on this architecture
|
||||
// so the tests exercise them directly, independent of what the public
|
||||
// Checksum dispatches to on the running CPU. Without this, running the
|
||||
// suite on a non-AVX2 machine compared gvisor against itself and left the
|
||||
// assembly untested — silently. available=false makes the test skip loudly
|
||||
// instead.
|
||||
var archImpls = []archImpl{
|
||||
{name: "avx2", fn: checksumAVX2, available: hasAVX2},
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
package checksum
|
||||
|
||||
// archImpls exposes every hand-written implementation on this architecture
|
||||
// for direct testing; see export_amd64_test.go for the rationale. NEON is
|
||||
// mandatory in armv8, so it is always available.
|
||||
var archImpls = []archImpl{
|
||||
{name: "neon", fn: checksumNEON, available: true},
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
//go:build !amd64 && !arm64
|
||||
|
||||
package checksum
|
||||
|
||||
// No hand-written implementations on this architecture; the dispatcher is
|
||||
// pure gvisor and there is nothing separate to test.
|
||||
var archImpls []archImpl
|
||||
+13
-3
@@ -4,15 +4,25 @@ import (
|
||||
"io"
|
||||
"net/netip"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
)
|
||||
|
||||
// defaultBatchBufSize is the per-Queue scratch size for Read on backends
|
||||
// that don't do TSO segmentation. 65535 covers any single IP packet.
|
||||
const defaultBatchBufSize = 65535
|
||||
|
||||
type Device interface {
|
||||
io.ReadWriteCloser
|
||||
io.Closer
|
||||
Activate() error
|
||||
Networks() []netip.Prefix
|
||||
Name() string
|
||||
RoutesFor(netip.Addr) routing.Gateways
|
||||
SupportsMultiqueue() bool
|
||||
NewMultiQueueReader() (io.ReadWriteCloser, error)
|
||||
// Queues returns the device's packet queues, opening additional ones as
|
||||
// needed until there are n. Platforms without multiqueue support return
|
||||
// their single queue regardless of n, so callers must size reader loops
|
||||
// to len(result), not n; implementations never return more than n. An
|
||||
// error means a queue that should have opened could not; the caller owns
|
||||
// cleanup via Close. Called once, during interface activation.
|
||||
Queues(n int) ([]tio.Queue, error)
|
||||
}
|
||||
|
||||
@@ -3,10 +3,9 @@
|
||||
package overlaytest
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io"
|
||||
"net/netip"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
)
|
||||
|
||||
@@ -31,20 +30,16 @@ func (NoopTun) Name() string {
|
||||
return "noop"
|
||||
}
|
||||
|
||||
func (NoopTun) Read([]byte) (int, error) {
|
||||
return 0, nil
|
||||
func (NoopTun) Read() ([]tio.Packet, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (NoopTun) Write([]byte) (int, error) {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
func (NoopTun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (NoopTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, errors.New("unsupported")
|
||||
func (NoopTun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{NoopTun{}}, nil
|
||||
}
|
||||
|
||||
func (NoopTun) Close() error {
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"os"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// blockOn parks the calling goroutine until fd is ready or shutdownFd signals teardown.
|
||||
// (events is POLLIN for reads, POLLOUT for writes)
|
||||
// It builds the pollfd array on the stack every call, so concurrent callers on the same Queue never share Revents storage.
|
||||
//
|
||||
// Returns os.ErrClosed when shutdown was signaled (POLLIN on shutdownFd)
|
||||
// or either fd reported a problem condition (POLLHUP|POLLNVAL|POLLERR).
|
||||
func blockOn(fd, shutdownFd int32, events int16) error {
|
||||
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
||||
pfds := [2]unix.PollFd{
|
||||
{Fd: fd, Events: events},
|
||||
{Fd: shutdownFd, Events: unix.POLLIN},
|
||||
}
|
||||
var err error
|
||||
for {
|
||||
_, err = unix.Poll(pfds[:], -1)
|
||||
if err != unix.EINTR {
|
||||
break
|
||||
}
|
||||
}
|
||||
tunEvents := pfds[0].Revents
|
||||
shutdownEvents := pfds[1].Revents
|
||||
// Check err before trusting the potentially bogus bits we just got.
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
||||
return os.ErrClosed
|
||||
}
|
||||
if tunEvents&problemFlags != 0 {
|
||||
return os.ErrClosed
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"sync/atomic"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
type offloadQueueSet struct {
|
||||
pq []*Offload
|
||||
// pqi is exactly the same as pq, but stored as the interface type
|
||||
pqi []Queue
|
||||
shutdownFd int
|
||||
// usoEnabled is true when newTun successfully negotiated TUN_F_USO4|6 with the kernel.
|
||||
// Queues created by Add inherit this and surface it via Offload.USOSupported so coalescers can gate USO emission.
|
||||
usoEnabled bool
|
||||
closed atomic.Bool
|
||||
// l is handed to each queue for its bad-vnet-header drop logging.
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
// NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do TSO segmentation.
|
||||
// usoEnabled tells downstream queues whether the kernel agreed to deliver/accept GSO_UDP_L4 superpackets.
|
||||
func NewOffloadQueueSet(usoEnabled bool, l *slog.Logger) (QueueSet, error) {
|
||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||
}
|
||||
|
||||
out := &offloadQueueSet{
|
||||
pq: []*Offload{},
|
||||
pqi: []Queue{},
|
||||
shutdownFd: shutdownFd,
|
||||
usoEnabled: usoEnabled,
|
||||
l: l,
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *offloadQueueSet) Queues() []Queue {
|
||||
return c.pqi
|
||||
}
|
||||
|
||||
func (c *offloadQueueSet) Add(fd int) error {
|
||||
if c.closed.Load() {
|
||||
return errors.New("queue set already closed")
|
||||
}
|
||||
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled, c.l)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.pq = append(c.pq, x)
|
||||
c.pqi = append(c.pqi, x)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *offloadQueueSet) wakeForShutdown() error {
|
||||
var buf [8]byte
|
||||
binary.NativeEndian.PutUint64(buf[:], 1)
|
||||
_, err := unix.Write(c.shutdownFd, buf[:])
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *offloadQueueSet) Close() error {
|
||||
if c.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
|
||||
errs := []error{}
|
||||
|
||||
// Signal all readers blocked in poll to wake up and exit.
|
||||
// They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
|
||||
if err := c.wakeForShutdown(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
|
||||
// Close the per-queue tun fds; this also unblocks any in-flight reads.
|
||||
for _, x := range c.pq {
|
||||
if err := x.Close(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Close the shutdown eventfd last: every reader's pollfd set references it,
|
||||
// so it must outlive the wake + per-queue teardown above.
|
||||
if err := unix.Close(c.shutdownFd); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
c.shutdownFd = -1
|
||||
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync/atomic"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
type pollQueueSet struct {
|
||||
pq []*Poll
|
||||
// pqi is exactly the same as pq, but stored as the interface type
|
||||
pqi []Queue
|
||||
shutdownFd int
|
||||
closed atomic.Bool
|
||||
}
|
||||
|
||||
func NewPollQueueSet() (QueueSet, error) {
|
||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||
}
|
||||
|
||||
out := &pollQueueSet{
|
||||
pq: []*Poll{},
|
||||
pqi: []Queue{},
|
||||
shutdownFd: shutdownFd,
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *pollQueueSet) Queues() []Queue {
|
||||
return c.pqi
|
||||
}
|
||||
|
||||
func (c *pollQueueSet) Add(fd int) error {
|
||||
if c.closed.Load() {
|
||||
return errors.New("queue set already closed")
|
||||
}
|
||||
x, err := newPoll(fd, c.shutdownFd)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.pq = append(c.pq, x)
|
||||
c.pqi = append(c.pqi, x)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *pollQueueSet) wakeForShutdown() error {
|
||||
var buf [8]byte
|
||||
binary.NativeEndian.PutUint64(buf[:], 1)
|
||||
_, err := unix.Write(int(c.shutdownFd), buf[:])
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *pollQueueSet) Close() error {
|
||||
if c.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
|
||||
errs := []error{}
|
||||
|
||||
// Signal all readers blocked in poll to wake up and exit.
|
||||
// They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
|
||||
if err := c.wakeForShutdown(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
|
||||
// Close the per-queue tun fds; this also unblocks any in-flight reads.
|
||||
for _, x := range c.pq {
|
||||
if err := x.Close(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Close the shutdown eventfd last: every reader's pollfd set references it,
|
||||
// so it must outlive the wake + per-queue teardown above.
|
||||
if err := unix.Close(c.shutdownFd); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
c.shutdownFd = -1
|
||||
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
|
||||
package tio
|
||||
|
||||
import "testing"
|
||||
|
||||
// fakeBatch stands in for batch.TxBatcher inside the bench — same shape
|
||||
// of pointer-capturing closure that sendInsideMessage builds.
|
||||
type fakeBatch struct{ buf [65536]byte }
|
||||
|
||||
func (b *fakeBatch) Reserve(sz int) []byte { return b.buf[:sz] }
|
||||
func (b *fakeBatch) Commit([]byte) {}
|
||||
|
||||
type fakeHostInfo struct {
|
||||
remoteIndexId uint32
|
||||
counter uint64
|
||||
}
|
||||
type fakeIface struct {
|
||||
rebindCount uint8
|
||||
hi *fakeHostInfo
|
||||
}
|
||||
|
||||
// BenchmarkSegmentSuperpacketAllocsTSO measures allocation per
|
||||
// SegmentSuperpacket call when a closure captures pointer-bearing
|
||||
// receivers — the realistic shape of sendInsideMessage's closure.
|
||||
func BenchmarkSegmentSuperpacketAllocsTSO(b *testing.B) {
|
||||
const mss = 1400
|
||||
const numSeg = 32
|
||||
pkt := buildTSOv6(mss*numSeg, mss)
|
||||
gso := GSOInfo{
|
||||
Size: mss,
|
||||
HdrLen: 60, // 40 (IPv6) + 20 (TCP)
|
||||
CsumStart: 40,
|
||||
Proto: GSOProtoTCP,
|
||||
}
|
||||
p := Packet{Bytes: pkt, GSO: gso}
|
||||
|
||||
hi := &fakeHostInfo{remoteIndexId: 0xdeadbeef}
|
||||
f := &fakeIface{rebindCount: 7, hi: hi}
|
||||
fb := &fakeBatch{}
|
||||
|
||||
// SegmentSuperpacket consumes pkt destructively; refresh from a master
|
||||
// copy each iter (matches the production pattern where every TUN read
|
||||
// hands the segmenter a fresh kernel-supplied buffer).
|
||||
master := append([]byte(nil), pkt...)
|
||||
work := make([]byte, len(pkt))
|
||||
p.Bytes = work
|
||||
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
copy(work, master)
|
||||
err := SegmentSuperpacket(p, func(seg []byte) error {
|
||||
out := fb.Reserve(16 + len(seg) + 16)
|
||||
out[0] = byte(f.rebindCount)
|
||||
out[1] = byte(hi.counter)
|
||||
hi.counter++
|
||||
fb.Commit(out)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
b.Fatalf("SegmentSuperpacket: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
//go:build !linux || android
|
||||
|
||||
package tio
|
||||
|
||||
import "fmt"
|
||||
|
||||
func protoFromGSOType(_ uint8) (GSOProto, error) {
|
||||
return 0, fmt.Errorf("GSO unsupported")
|
||||
}
|
||||
|
||||
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
||||
if pkt.GSO.IsSuperpacket() {
|
||||
return fmt.Errorf("tio: GSO superpacket on platform without segmentation support")
|
||||
}
|
||||
return fn(pkt.Bytes)
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
package tio
|
||||
|
||||
import "io"
|
||||
|
||||
// singleQueue adapts a legacy one-datagram-per-Read source into a Queue.
|
||||
// Read fills a private scratch buffer and returns exactly one Packet whose
|
||||
// Bytes borrow from that buffer, valid only until the next Read, per the Queue contract.
|
||||
// Single-reader like every Queue; Write is exactly as safe for concurrent use as the underlying source's Write.
|
||||
type singleQueue struct {
|
||||
rw io.ReadWriter
|
||||
closer io.Closer // nil: Close is a no-op (the source is shared and owned elsewhere)
|
||||
buf []byte
|
||||
ret [1]Packet
|
||||
}
|
||||
|
||||
// NewSingleQueue wraps a one-datagram-per-Read ReadWriteCloser (a legacy tun device) into a Queue.
|
||||
// bufSize is the per-queue read scratch size and must be at least the largest datagram the source can return.
|
||||
// Close closes rwc.
|
||||
func NewSingleQueue(rwc io.ReadWriteCloser, bufSize int) Queue {
|
||||
return &singleQueue{rw: rwc, closer: rwc, buf: make([]byte, bufSize)}
|
||||
}
|
||||
|
||||
// NewSingleQueueNoClose is NewSingleQueue for a source owned by someone else,
|
||||
// e.g. several queues sharing one device. Close on the returned Queue is a
|
||||
// no-op so one queue can't tear the shared source out from under its
|
||||
// siblings; the owner remains responsible for closing the source itself.
|
||||
func NewSingleQueueNoClose(rw io.ReadWriter, bufSize int) Queue {
|
||||
return &singleQueue{rw: rw, buf: make([]byte, bufSize)}
|
||||
}
|
||||
|
||||
func (q *singleQueue) Read() ([]Packet, error) {
|
||||
n, err := q.rw.Read(q.buf)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
q.ret[0] = Packet{Bytes: q.buf[:n]}
|
||||
return q.ret[:], nil
|
||||
}
|
||||
|
||||
func (q *singleQueue) Write(p []byte) (int, error) {
|
||||
return q.rw.Write(p)
|
||||
}
|
||||
|
||||
func (q *singleQueue) Close() error {
|
||||
if q.closer == nil {
|
||||
return nil
|
||||
}
|
||||
return q.closer.Close()
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
package tio
|
||||
|
||||
import (
|
||||
"io"
|
||||
)
|
||||
|
||||
// QueueSet holds one or many Queue objects and helps close them in an orderly way.
|
||||
type QueueSet interface {
|
||||
io.Closer
|
||||
Queues() []Queue
|
||||
|
||||
// Add takes a tun fd, adds it to the set, and prepares it for use as a Queue.
|
||||
Add(fd int) error
|
||||
}
|
||||
|
||||
// Capabilities advertises which kernel offload features a Queue successfully negotiated.
|
||||
// Callers consult this to decide which coalescers to wire onto the write path.
|
||||
type Capabilities struct {
|
||||
// TSO means the FD was opened with IFF_VNET_HDR and the kernel agreed to TUN_F_TSO4|TSO6,
|
||||
// and WriteGSO with GSOProtoTCP is safe.
|
||||
TSO bool
|
||||
// USO means the kernel additionally agreed to TUN_F_USO4|USO6,
|
||||
// so WriteGSO with GSOProtoUDP is safe. Linux ≥ 6.2.
|
||||
USO bool
|
||||
}
|
||||
|
||||
// Queue is a readable/writable Poll queue.
|
||||
// Concurrency contract: a single read goroutine drives Read; plain Write is safe for concurrent callers;
|
||||
// WriteGSO (on Queues that implement GSOWriter) is single-writer per queue.
|
||||
//
|
||||
// Close on an individual Queue does NOT unblock a Read parked in poll — closing an fd
|
||||
// never wakes its pollers. Orderly teardown goes through the owning QueueSet's Close,
|
||||
// which first signals a shared shutdown eventfd every reader polls alongside its own fd.
|
||||
// That eventfd is a set-wide kill switch: once signaled, every Queue in the set returns
|
||||
// os.ErrClosed from Read, so it cannot be used to stop a single Queue.
|
||||
type Queue interface {
|
||||
io.Closer
|
||||
|
||||
// Read returns one or more packets.
|
||||
// The returned Packet.Bytes slices are borrowed from the Queue's internal buffer and are only valid
|
||||
// until the next Read or Close on this Queue.
|
||||
// A Packet may carry a GSO/USO superpacket (see GSOInfo)
|
||||
// Single-reader only: not safe for concurrent Reads (it reuses per-queue rx scratch each call).
|
||||
Read() ([]Packet, error)
|
||||
|
||||
// Write emits a single packet on the plaintext (outside→inside) delivery path.
|
||||
// Safe for concurrent use.
|
||||
Write(p []byte) (int, error)
|
||||
}
|
||||
|
||||
// Packet is the unit Queue.Read returns.
|
||||
// Bytes points into the queue's internal buffer and is only valid until the next Read or Close on the queue that produced it.
|
||||
// GSO is the zero value for an already-segmented IP datagram;
|
||||
// when non-zero it describes a kernel-supplied TSO/USO superpacket the caller must segment before consuming.
|
||||
type Packet struct {
|
||||
Bytes []byte
|
||||
GSO GSOInfo
|
||||
}
|
||||
|
||||
// GSOInfo describes a kernel-supplied superpacket sitting in Packet.Bytes.
|
||||
// The zero value means Bytes is one regular IP datagram and no segmentation is required.
|
||||
type GSOInfo struct {
|
||||
// Size is the GSO segment size: max payload bytes per segment
|
||||
// (== TCP MSS for TSO, == UDP payload chunk for USO). Zero means not a superpacket.
|
||||
Size uint16
|
||||
// HdrLen is the total L3+L4 header length within Bytes (already corrected via correctHdrLen, so safe to slice on).
|
||||
HdrLen uint16
|
||||
// CsumStart is the L4 header offset inside Bytes (== L3 header length).
|
||||
CsumStart uint16
|
||||
// Proto picks the L4 protocol (TCP or UDP) so the segmenter knows which checksum/header layout to apply.
|
||||
Proto GSOProto
|
||||
}
|
||||
|
||||
// IsSuperpacket reports whether g describes a multi-segment GSO/USO
|
||||
// superpacket that needs segmentation before its bytes can be encrypted and sent on the wire.
|
||||
func (g GSOInfo) IsSuperpacket() bool { return g.Size > 0 }
|
||||
|
||||
// Clone returns a Packet whose Bytes is a freshly allocated copy of p.Bytes,
|
||||
// safe to retain past the next Read or Close on the originating Queue.
|
||||
// GSO metadata is copied verbatim.
|
||||
// Use this only when a caller needs the data to outlive the borrowed-slice contract.
|
||||
func (p Packet) Clone() Packet {
|
||||
if p.Bytes == nil {
|
||||
return p
|
||||
}
|
||||
cp := make([]byte, len(p.Bytes))
|
||||
copy(cp, p.Bytes)
|
||||
return Packet{Bytes: cp, GSO: p.GSO}
|
||||
}
|
||||
|
||||
// CapsProvider is an optional interface implemented by Queues that negotiate kernel offload features at open time.
|
||||
// Callers pick a write-path coalescer based on the result.
|
||||
// Queues that don't implement it are treated as having no offload capability.
|
||||
type CapsProvider interface {
|
||||
Capabilities() Capabilities
|
||||
}
|
||||
|
||||
// GSOProto selects the L4 protocol for a GSO superpacket.
|
||||
// Determines which VIRTIO_NET_HDR_GSO_* type the writer stamps and which checksum offset
|
||||
// inside the transport header virtio NEEDS_CSUM expects.
|
||||
type GSOProto uint8
|
||||
|
||||
const (
|
||||
GSOProtoUnknown GSOProto = iota
|
||||
GSOProtoTCP
|
||||
GSOProtoUDP
|
||||
)
|
||||
|
||||
// GSOWriter is implemented by Queues that can emit a TCP or UDP superpacket
|
||||
// assembled from a header prefix plus one or more borrowed payload fragments,
|
||||
// in a single vectored write (writev with a leading virtio_net_hdr).
|
||||
// This lets the coalescer avoid copying payload bytes between the caller's decrypt buffer and the TUN.
|
||||
// Backends without GSO support do not implement this interface and coalescing is skipped.
|
||||
//
|
||||
// hdr contains the IPv4/IPv6 header prefix (mutable: callers will have filled in total length and IP csum).
|
||||
// transportHdr is the TCP or UDP header
|
||||
// (mutable: the L4 checksum field must hold the pseudo-header partial, single-fold not inverted, per virtio NEEDS_CSUM semantics).
|
||||
// pays are non-overlapping payload fragments whose concatenation is the full superpacket payload.
|
||||
// They are read-only from the writer's perspective and must remain valid until the call returns.
|
||||
// Every segment in pays except possibly the last must be exactly the same size.
|
||||
// proto picks the L4 protocol so the writer knows which gsoType / CsumOffset to set.
|
||||
//
|
||||
// Callers should also consult CapsProvider (via SupportsGSO) for the per-protocol negotiated capability:
|
||||
// USO may not have been negotiated even when TSO was.
|
||||
type GSOWriter interface {
|
||||
io.Writer
|
||||
CapsProvider
|
||||
WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error
|
||||
}
|
||||
|
||||
// SupportsGSO reports whether w implements GSOWriter and the underlying
|
||||
// queue advertises the negotiated capability for `want`.
|
||||
func SupportsGSO(w io.Writer, want GSOProto) (GSOWriter, bool) {
|
||||
gw, ok := w.(GSOWriter)
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
caps := gw.Capabilities()
|
||||
switch want {
|
||||
case GSOProtoTCP:
|
||||
return gw, caps.TSO
|
||||
case GSOProtoUDP:
|
||||
return gw, caps.USO
|
||||
default:
|
||||
return gw, false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,417 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"os"
|
||||
"sync/atomic"
|
||||
"syscall"
|
||||
"unsafe"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio/virtio"
|
||||
)
|
||||
|
||||
const maxSuperpacketLen = 65535
|
||||
|
||||
// tunRxBufSize is the per-Read worst-case footprint inside rxBuf: one kernel-supplied packet body, which is at most ~64 KiB.
|
||||
// Segmentation happens at encrypt time on a per-routine MTU-sized scratch
|
||||
// (see SegmentSuperpacket), so rxBuf only holds raw kernel-supplied bytes.
|
||||
// We round up to give margin for the drain headroom check below.
|
||||
const tunRxBufSize = 64 * 1024
|
||||
|
||||
// tunRxBufCap is the total size we allocate for the per-reader rx buffer.
|
||||
// Each drain iteration consumes up to tunRxBufSize of headroom for the kernel-supplied bytes.
|
||||
// Sized to eight such iterations so a single poll wake can drain several TSO/USO superpackets under bulk load,
|
||||
// amortizing the wake and giving the sendmmsg planner longer same-destination runs.
|
||||
// Hold latency stays bounded because listenIn flushes its send batch incrementally rather than only at end-of-drain.
|
||||
const tunRxBufCap = tunRxBufSize * 8
|
||||
|
||||
// tunDrainCap caps how many packets a single Read will accumulate via the post-wake drain loop.
|
||||
// Sized to soak up a burst of small ACKs while bounding how much work a single caller holds before handing off.
|
||||
const tunDrainCap = 64
|
||||
|
||||
// gsoMaxIovs caps the iovec budget WriteGSO assembles per call:
|
||||
// 3 fixed entries (virtio_net_hdr, IP hdr, transport hdr), plus up to gsoMaxIovs-3 payload fragments.
|
||||
// Sized comfortably above the typical kernel GSO segment cap (Linux UDP_GRO is 64)
|
||||
// so realistic coalesced bursts never touch the limit.
|
||||
// iovecs are tiny (16 bytes), so the entire scratch is 4 KiB.
|
||||
// WriteGSO returns an error rather than reallocating when a caller exceeds this budget.
|
||||
const gsoMaxIovs = 256
|
||||
|
||||
// validVnetHdr is the 10-byte virtio_net_hdr we prepend to every non-GSO TUN write.
|
||||
// Only flag set is VIRTIO_NET_HDR_F_DATA_VALID. Note the tun write path
|
||||
// (__virtio_net_hdr_to_skb) ignores this bit — only the virtio-net driver's RX
|
||||
// helper honors it — so packets land CHECKSUM_NONE and the stack verifies the
|
||||
// L4 checksum anyway. What matters here is what the header does NOT say:
|
||||
// no NEEDS_CSUM, so the kernel is never asked to finish a checksum.
|
||||
// All packets that reach the plain Write paths already carry a valid L4 checksum.
|
||||
var validVnetHdr = [virtio.Size]byte{unix.VIRTIO_NET_HDR_F_DATA_VALID}
|
||||
|
||||
// Offload wraps a TUN file descriptor with poll-based reads. The FD provided will be changed to non-blocking.
|
||||
// A shared eventfd allows Close to wake all readers blocked in poll.
|
||||
//
|
||||
// Field order is deliberate: the read-mostly fds and the writer-owned GSO scratch fill
|
||||
// the first cache line, and the state the reader mutates per packet (rxOff, pending,
|
||||
// readIovs) all sits after it, so per-packet reader stores never invalidate the line
|
||||
// concurrent Write callers load fd from.
|
||||
type Offload struct {
|
||||
fd int
|
||||
shutdownFd int
|
||||
// usoEnabled records whether the kernel agreed to TUN_F_USO* on this FD,
|
||||
// so writers can decide whether emitting GSO_UDP_L4 superpackets is safe.
|
||||
usoEnabled bool
|
||||
closed atomic.Bool
|
||||
|
||||
// gsoHdrBuf is a per-queue 10-byte scratch for the virtio_net_hdr emitted
|
||||
// by WriteGSO. Kept separate from the read-only package-level validVnetHdr
|
||||
// so non-GSO Writes can ship that constant directly while WriteGSO
|
||||
// rewrites this scratch on every call.
|
||||
gsoHdrBuf [virtio.Size]byte
|
||||
// gsoIovs is the writev iovec scratch for WriteGSO. Pre-sized to
|
||||
// gsoMaxIovs at construction; never grown. WriteGSO returns an error
|
||||
// (and drops the call) if a caller hands it more fragments than fit.
|
||||
gsoIovs []unix.Iovec
|
||||
|
||||
rxBuf []byte // backing store for kernel-handed packets read this drain
|
||||
rxOff int // cursor into rxBuf for the current Read drain
|
||||
pending []Packet // packets returned from the most recent Read
|
||||
|
||||
// readVnetScratch holds the 10-byte virtio_net_hdr split off the front of
|
||||
// every TUN read via readv(2). Decoupling the header from the packet body
|
||||
// lets us read the body directly into rxBuf at the current rxOff with
|
||||
// no userspace copy on the GSO_NONE fast path.
|
||||
readVnetScratch [virtio.Size]byte
|
||||
// readIovs is the readv(2) iovec scratch wired once at construction,
|
||||
// iovec[0] points at readVnetScratch
|
||||
// iovec[1].Base/Len is updated per read to address the current rxBuf slot.
|
||||
readIovs [2]unix.Iovec
|
||||
|
||||
// l is only consulted on the rare bad-vnet-header drop path; it lives
|
||||
// after the hot state on purpose. May be nil (tests); drops go unlogged then.
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
func newOffload(fd int, shutdownFd int, usoEnabled bool, l *slog.Logger) (*Offload, error) {
|
||||
if err := unix.SetNonblock(fd, true); err != nil {
|
||||
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
||||
}
|
||||
|
||||
out := &Offload{
|
||||
fd: fd,
|
||||
shutdownFd: shutdownFd,
|
||||
usoEnabled: usoEnabled,
|
||||
closed: atomic.Bool{},
|
||||
l: l,
|
||||
|
||||
rxBuf: make([]byte, tunRxBufCap),
|
||||
gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs),
|
||||
}
|
||||
|
||||
out.gsoIovs[0].Base = &out.gsoHdrBuf[0]
|
||||
out.gsoIovs[0].SetLen(virtio.Size)
|
||||
|
||||
// readIovs[0] is wired once to the virtio_net_hdr scratch; per-read we
|
||||
// only repoint readIovs[1] at the next rxBuf slot (see readPacket).
|
||||
out.readIovs[0].Base = &out.readVnetScratch[0]
|
||||
out.readIovs[0].SetLen(virtio.Size)
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (r *Offload) blockOnRead() error {
|
||||
return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLIN)
|
||||
}
|
||||
|
||||
func (r *Offload) blockOnWrite() error {
|
||||
return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLOUT)
|
||||
}
|
||||
|
||||
// readPacket issues a single readv(2), splitting the virtio_net_hdr off into readVnetScratch
|
||||
// and reading the packet body directly into rxBuf at the current rxOff.
|
||||
// Returns the body length (zero virtio header bytes, just the IP packet/superpacket).
|
||||
// block controls whether EAGAIN is retried via poll: the initial read of a drain blocks; subsequent drain reads do not.
|
||||
func (r *Offload) readPacket(block bool) (int, error) {
|
||||
for {
|
||||
r.readIovs[1].Base = &r.rxBuf[r.rxOff]
|
||||
r.readIovs[1].SetLen(len(r.rxBuf) - r.rxOff)
|
||||
n, _, errno := syscall.Syscall(unix.SYS_READV, uintptr(r.fd), uintptr(unsafe.Pointer(&r.readIovs[0])), uintptr(len(r.readIovs)))
|
||||
if errno == 0 {
|
||||
if int(n) < virtio.Size {
|
||||
return 0, fmt.Errorf("tun read shorter than virtio_net_hdr: %d bytes", n)
|
||||
}
|
||||
return int(n) - virtio.Size, nil
|
||||
}
|
||||
if errno == unix.EAGAIN {
|
||||
if !block {
|
||||
return 0, errno
|
||||
}
|
||||
if err := r.blockOnRead(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
continue
|
||||
}
|
||||
if errno == unix.EINTR {
|
||||
continue
|
||||
}
|
||||
if errno == unix.EBADF {
|
||||
return 0, os.ErrClosed
|
||||
}
|
||||
return 0, errno
|
||||
}
|
||||
}
|
||||
|
||||
// Read returns one or more packets from the tun.
|
||||
// Each Packet either carries a single ready-to-use IP datagram (GSO zero) or a TSO/USO superpacket plus the GSOInfo a caller needs to segment it (see SegmentSuperpacket).
|
||||
// The first read blocks via poll; once the fd is known readable we drain additional packets non-blocking until:
|
||||
// - the kernel queue is empty (EAGAIN)
|
||||
// - we've collected tunDrainCap packets,
|
||||
// - or we're out of rxBuf headroom.
|
||||
//
|
||||
// This amortizes the poll wake over bursts of small packets (e.g. TCP ACKs).
|
||||
// Packet.Bytes slices point into the Offload's internal buffer and are only valid until the next Read or Close on this Queue.
|
||||
func (r *Offload) Read() ([]Packet, error) {
|
||||
r.pending = r.pending[:0]
|
||||
r.rxOff = 0
|
||||
|
||||
// Initial (blocking) read.
|
||||
// Retry on decode errors so a single bad packet does not stall the reader.
|
||||
for {
|
||||
n, err := r.readPacket(true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := r.decodeRead(n); err != nil {
|
||||
// Drop and read again. A bad packet should not kill the reader,
|
||||
// but a systematic decode failure must not be invisible either.
|
||||
r.logDroppedRead(err)
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
|
||||
// Drain: non-blocking reads until the kernel queue is empty, the drain
|
||||
// cap is reached, or rxBuf no longer has room for another worst-case
|
||||
// kernel-supplied packet (tunRxBufSize).
|
||||
for len(r.pending) < tunDrainCap && tunRxBufCap-r.rxOff >= tunRxBufSize {
|
||||
n, err := r.readPacket(false)
|
||||
if err != nil {
|
||||
// EAGAIN / EINTR / anything else: stop draining. We already
|
||||
// have a valid batch from the first read.
|
||||
break
|
||||
}
|
||||
if n <= 0 {
|
||||
break
|
||||
}
|
||||
if err := r.decodeRead(n); err != nil {
|
||||
// Drop this packet and stop the drain; we'd rather hand off
|
||||
// what we have than keep spinning here.
|
||||
r.logDroppedRead(err)
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
return r.pending, nil
|
||||
}
|
||||
|
||||
// logDroppedRead reports a tun packet dropped for a bad/unsupported virtio
|
||||
// header. Debug-gated so the happy path never pays for attribute assembly.
|
||||
func (r *Offload) logDroppedRead(err error) {
|
||||
if r.l != nil && r.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
r.l.Debug("dropping tun packet with bad virtio header", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// decodeRead processes the packet sitting in rxBuf at rxOff (length pktLen).
|
||||
// The bytes stay in rxBuf:
|
||||
// - for GSO_NONE we slice them as a regular IP datagram (running finishChecksum if NEEDS_CSUM is set);
|
||||
// - for TSO/USO superpackets we attach the corrected GSO metadata, so the caller can segment lazily at encrypt time.
|
||||
//
|
||||
// rxOff advances by pktLen on success
|
||||
func (r *Offload) decodeRead(pktLen int) error {
|
||||
if pktLen <= 0 {
|
||||
return fmt.Errorf("short tun read: %d", pktLen)
|
||||
}
|
||||
var hdr virtio.Hdr
|
||||
hdr.Decode(r.readVnetScratch[:])
|
||||
|
||||
body := r.rxBuf[r.rxOff : r.rxOff+pktLen]
|
||||
|
||||
if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
|
||||
if err := virtio.FinishChecksum(body, hdr); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
r.pending = append(r.pending, Packet{Bytes: body})
|
||||
r.rxOff += pktLen
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := virtio.CheckValid(body, hdr); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := virtio.CorrectHdrLen(body, &hdr); err != nil {
|
||||
return err
|
||||
}
|
||||
proto, err := protoFromGSOType(hdr.GSOType())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
r.pending = append(r.pending, Packet{
|
||||
Bytes: body,
|
||||
GSO: GSOInfo{
|
||||
Size: hdr.GSOSize,
|
||||
HdrLen: hdr.HdrLen,
|
||||
CsumStart: hdr.CsumStart,
|
||||
Proto: proto,
|
||||
},
|
||||
})
|
||||
r.rxOff += pktLen
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Offload) Write(buf []byte) (int, error) {
|
||||
if len(buf) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
iovs := [2]unix.Iovec{
|
||||
{Base: &validVnetHdr[0]},
|
||||
{Base: &buf[0]},
|
||||
}
|
||||
iovs[0].SetLen(virtio.Size)
|
||||
iovs[1].SetLen(len(buf))
|
||||
return r.rawWrite(unsafe.Slice(&iovs[0], 2))
|
||||
}
|
||||
|
||||
func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
|
||||
for {
|
||||
n, _, errno := syscall.Syscall(unix.SYS_WRITEV, uintptr(r.fd), uintptr(unsafe.Pointer(&iovs[0])), uintptr(len(iovs)))
|
||||
if errno == 0 {
|
||||
if int(n) < virtio.Size {
|
||||
return 0, io.ErrShortWrite
|
||||
}
|
||||
return int(n) - virtio.Size, nil
|
||||
}
|
||||
if errno == unix.EAGAIN {
|
||||
if err := r.blockOnWrite(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
continue
|
||||
}
|
||||
if errno == unix.EINTR {
|
||||
continue
|
||||
}
|
||||
if errno == unix.EBADF {
|
||||
return 0, os.ErrClosed
|
||||
}
|
||||
return 0, errno
|
||||
}
|
||||
}
|
||||
|
||||
// Capabilities reports the offload features negotiated for this Queue. TSO
|
||||
// is always true for Offload (we only construct it on IFF_VNET_HDR FDs);
|
||||
// USO is true only when the kernel agreed to TUN_F_USO4|6 at open time (Linux ≥ 6.2).
|
||||
func (r *Offload) Capabilities() Capabilities {
|
||||
return Capabilities{TSO: true, USO: r.usoEnabled}
|
||||
}
|
||||
|
||||
func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error {
|
||||
if len(pays) == 0 {
|
||||
// There are no payload fragments. There is nothing to send.
|
||||
return nil
|
||||
}
|
||||
var csumOff uint16 // csumOff is the offset of the L4 checksum field in transportHdr
|
||||
switch proto {
|
||||
case GSOProtoUDP:
|
||||
csumOff = 6
|
||||
case GSOProtoTCP:
|
||||
csumOff = 16
|
||||
default:
|
||||
return fmt.Errorf("unknown GSO proto: %d", proto)
|
||||
}
|
||||
// Incorrect geometry must cause an error, not a silent drop.
|
||||
// No sane packet should ever make it inside this branch.
|
||||
if len(hdr) == 0 || len(transportHdr) < int(csumOff)+2 {
|
||||
return fmt.Errorf("tio: WriteGSO header too short: ip=%d transport=%d (csum field at %d)", len(hdr), len(transportHdr), csumOff)
|
||||
}
|
||||
// Make the iovec array: [virtio_hdr, hdr, transportHdr, pays...].
|
||||
// The constructor attaches r.gsoIovs[0] to gsoHdrBuf. That entry does not change.
|
||||
need := 3 + len(pays)
|
||||
if need > cap(r.gsoIovs) {
|
||||
return fmt.Errorf("tio: WriteGSO needs %d iovecs but cap is %d", need, cap(r.gsoIovs))
|
||||
}
|
||||
r.gsoIovs = r.gsoIovs[:need]
|
||||
r.gsoIovs[1].Base = &hdr[0]
|
||||
r.gsoIovs[1].SetLen(len(hdr))
|
||||
r.gsoIovs[2].Base = &transportHdr[0]
|
||||
r.gsoIovs[2].SetLen(len(transportHdr))
|
||||
|
||||
segSize := len(pays[0])
|
||||
total := len(hdr) + len(transportHdr)
|
||||
for i, p := range pays {
|
||||
if len(p) == 0 {
|
||||
// The coalescers route zero-payload packets down the non-GSO path,
|
||||
// so an empty fragment means the caller's accounting is broken.
|
||||
return fmt.Errorf("tio: WriteGSO empty payload fragment %d of %d", i, len(pays))
|
||||
} else if len(p) > segSize || (len(p) < segSize && i != len(pays)-1) {
|
||||
// all segments must be the same size, except for the last one
|
||||
return fmt.Errorf("tio: WriteGSO fragment %d is %dB, want %dB segments (only the last may be shorter)", i, len(p), segSize)
|
||||
}
|
||||
total += len(p)
|
||||
r.gsoIovs[3+i].Base = &p[0]
|
||||
r.gsoIovs[3+i].SetLen(len(p))
|
||||
}
|
||||
// This check keeps `total` in the uint16 range. Anything larger would wrap around and cause trouble.
|
||||
if total > maxSuperpacketLen {
|
||||
return fmt.Errorf("tio: WriteGSO superpacket %dB exceeds %d", total, maxSuperpacketLen)
|
||||
}
|
||||
|
||||
// A single segment ships as a plain checksummed packet (GSO_NONE, size 0).
|
||||
// Multiple segments carry the real GSO type and segSize, which the loop
|
||||
// above verified is the size of every fragment except possibly the last.
|
||||
gsoType := uint8(unix.VIRTIO_NET_HDR_GSO_NONE)
|
||||
if len(pays) > 1 {
|
||||
gsoType = gsoTypeFromProto(proto, hdr[0]>>4)
|
||||
if gsoType == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||
// gsoTypeFromProto only yields GSO_NONE for a bogus IP version nibble.
|
||||
// A multi-segment superpacket must carry a real GSO type, or the kernel would deliver it as a single jumbo packet.
|
||||
return fmt.Errorf("tio: WriteGSO IP version %d is not GSO-capable", hdr[0]>>4)
|
||||
}
|
||||
}
|
||||
var gsoSize uint16
|
||||
if gsoType != unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||
gsoSize = uint16(segSize)
|
||||
}
|
||||
virtio.EncodeHeader(
|
||||
r.gsoHdrBuf[:],
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
gsoType, /*gsoType*/
|
||||
uint16(len(hdr)+len(transportHdr)), /*hdrLen*/
|
||||
gsoSize, /*gsoSize*/
|
||||
uint16(len(hdr)), /*csumStart*/
|
||||
csumOff, /*csumOffset*/
|
||||
)
|
||||
|
||||
_, err := r.rawWrite(r.gsoIovs)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *Offload) Close() error {
|
||||
if r.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
|
||||
// shutdownFd is owned by the container, so we should not close it
|
||||
// Close the underlying fd but do NOT null r.fd: a reader may still be loading it in readPacket, and mutating the field would race that load.
|
||||
// That reader gets EBADF -> os.ErrClosed on its next syscall. A reader already parked in
|
||||
// poll is NOT woken by this close; only the QueueSet's shutdown eventfd wake does that (see Queue.Close docs).
|
||||
// closed.Swap already guarantees we only close once.
|
||||
return unix.Close(r.fd)
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"sync/atomic"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
type Poll struct {
|
||||
fd int
|
||||
shutdownFd int
|
||||
closed atomic.Bool
|
||||
|
||||
readBuf []byte
|
||||
batchRet [1]Packet
|
||||
}
|
||||
|
||||
// newPoll wraps an existing tun fd.
|
||||
// On failure it does NOT close fd: the caller owns fd and is the sole closer
|
||||
// (see pollQueueSet.Add callers in overlay/tun_linux.go, which unix.Close on Add error).
|
||||
// This matches the newOffload convention and keeps closes at exactly one on every path.
|
||||
func newPoll(fd int, shutdownFd int) (*Poll, error) {
|
||||
if err := unix.SetNonblock(fd, true); err != nil {
|
||||
return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err)
|
||||
}
|
||||
|
||||
out := &Poll{
|
||||
fd: fd,
|
||||
shutdownFd: shutdownFd,
|
||||
readBuf: make([]byte, 65535), // largest possible size Linux permits
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// blockOnRead waits until the Poll fd is readable or shutdown has been signaled.
|
||||
// Returns os.ErrClosed if Close was called.
|
||||
func (t *Poll) blockOnRead() error {
|
||||
return blockOn(int32(t.fd), int32(t.shutdownFd), unix.POLLIN)
|
||||
}
|
||||
|
||||
func (t *Poll) blockOnWrite() error {
|
||||
return blockOn(int32(t.fd), int32(t.shutdownFd), unix.POLLOUT)
|
||||
}
|
||||
|
||||
// TODO: port Offload's post-wake drain loop here so one poll wake amortizes
|
||||
// over a burst (up to tunDrainCap packets) instead of paying a syscall and a
|
||||
// wake per packet. Hosts on the TUNSETOFFLOAD-failure fallback or a tun.fd
|
||||
// config currently lose that batching. blockOn and the EAGAIN plumbing are
|
||||
// already shared; kept one-packet-per-Read for now to preserve behavior.
|
||||
func (t *Poll) Read() ([]Packet, error) {
|
||||
n, err := t.readOne(t.readBuf)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
t.batchRet[0] = Packet{Bytes: t.readBuf[:n]}
|
||||
return t.batchRet[:], nil
|
||||
}
|
||||
|
||||
func (t *Poll) readOne(to []byte) (int, error) {
|
||||
for {
|
||||
n, errno := unix.Read(t.fd, to)
|
||||
if errno == nil {
|
||||
return n, nil
|
||||
}
|
||||
switch errno {
|
||||
case unix.EAGAIN:
|
||||
if err := t.blockOnRead(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
case unix.EINTR:
|
||||
// retry
|
||||
case unix.EBADF:
|
||||
return 0, os.ErrClosed
|
||||
default:
|
||||
return 0, errno
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Write is safe for concurrent use
|
||||
func (t *Poll) Write(from []byte) (int, error) {
|
||||
for {
|
||||
n, errno := unix.Write(t.fd, from)
|
||||
if errno == nil {
|
||||
return n, nil
|
||||
}
|
||||
switch errno {
|
||||
case unix.EAGAIN:
|
||||
if err := t.blockOnWrite(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
case unix.EINTR:
|
||||
// retry
|
||||
case unix.EBADF:
|
||||
return 0, os.ErrClosed
|
||||
default:
|
||||
return 0, errno
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (t *Poll) Close() error {
|
||||
if t.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
|
||||
// shutdownFd is owned by the container, so we should not close it
|
||||
// Close the underlying fd but do NOT null t.fd: a reader may still be loading it in readOne, and mutating the field would race that load.
|
||||
// That reader gets EBADF -> os.ErrClosed on its next syscall. A reader already parked in
|
||||
// poll is NOT woken by this close; only the QueueSet's shutdown eventfd wake does that (see Queue.Close docs).
|
||||
// closed.Swap already guarantees we only close once.
|
||||
return unix.Close(t.fd)
|
||||
}
|
||||
@@ -0,0 +1,228 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
// +build linux,!android,!e2e_testing
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"log/slog"
|
||||
"os"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// newReadPipe returns a read fd. The matching write fd is registered for cleanup.
|
||||
// The caller takes ownership of the read fd (pass it into a QueueSet).
|
||||
func newReadPipe(t *testing.T) int {
|
||||
t.Helper()
|
||||
var fds [2]int
|
||||
if err := unix.Pipe2(fds[:], unix.O_CLOEXEC); err != nil {
|
||||
t.Fatalf("pipe2: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = unix.Close(fds[1]) })
|
||||
return fds[0]
|
||||
}
|
||||
|
||||
func TestPoll_WakeForShutdown_WakesFriends(t *testing.T) {
|
||||
pipe1 := newReadPipe(t)
|
||||
pipe2 := newReadPipe(t)
|
||||
parent, err := NewPollQueueSet()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, parent.Add(pipe1))
|
||||
require.NoError(t, parent.Add(pipe2))
|
||||
t.Cleanup(func() {
|
||||
_ = unix.Close(pipe1)
|
||||
_ = unix.Close(pipe2)
|
||||
})
|
||||
|
||||
readers := parent.Queues()
|
||||
errs := make([]error, len(readers))
|
||||
var wg sync.WaitGroup
|
||||
for i, r := range readers {
|
||||
wg.Add(1)
|
||||
go func(i int, r Queue) {
|
||||
defer wg.Done()
|
||||
_, errs[i] = r.Read()
|
||||
}(i, r)
|
||||
}
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
if err := parent.Close(); err != nil {
|
||||
t.Fatalf("Close: %v", err)
|
||||
}
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() { wg.Wait(); close(done) }()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("readers did not wake")
|
||||
}
|
||||
|
||||
for i, err := range errs {
|
||||
if !errors.Is(err, os.ErrClosed) {
|
||||
t.Errorf("reader %d: expected os.ErrClosed, got %v", i, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestPoll_ConcurrentWrite_NoRace hammers a single Poll queue from two writer
|
||||
// goroutines while a reader drains the other end of the pipe. The writers
|
||||
// overflow the pipe buffer, so both repeatedly park in blockOnWrite at the same
|
||||
// time — the exact scenario that raced on the old shared writePoll member
|
||||
// array. Run under -race; a shared-array regression trips the detector here.
|
||||
func TestPoll_ConcurrentWrite_NoRace(t *testing.T) {
|
||||
var fds [2]int
|
||||
require.NoError(t, unix.Pipe2(fds[:], unix.O_CLOEXEC))
|
||||
readFd, writeFd := fds[0], fds[1]
|
||||
|
||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = unix.Close(shutdownFd) })
|
||||
|
||||
p, err := newPoll(writeFd, shutdownFd)
|
||||
require.NoError(t, err)
|
||||
|
||||
const writers = 2
|
||||
const perWriter = 4000
|
||||
payload := make([]byte, 100)
|
||||
total := writers * perWriter * len(payload)
|
||||
|
||||
// Reader: drain the read end (blocking) until every writer's bytes are
|
||||
// consumed, so the writers keep making progress rather than wedging on a
|
||||
// permanently full pipe.
|
||||
readDone := make(chan struct{})
|
||||
go func() {
|
||||
defer close(readDone)
|
||||
buf := make([]byte, 4096)
|
||||
got := 0
|
||||
for got < total {
|
||||
n, rerr := unix.Read(readFd, buf)
|
||||
got += n
|
||||
if rerr != nil {
|
||||
if rerr == unix.EINTR {
|
||||
continue
|
||||
}
|
||||
return
|
||||
}
|
||||
if n == 0 { // EOF
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for w := 0; w < writers; w++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for i := 0; i < perWriter; i++ {
|
||||
if _, werr := p.Write(payload); werr != nil {
|
||||
t.Errorf("write: %v", werr)
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
select {
|
||||
case <-readDone:
|
||||
case <-time.After(10 * time.Second):
|
||||
t.Fatal("reader did not drain")
|
||||
}
|
||||
|
||||
require.NoError(t, p.Close())
|
||||
_ = unix.Close(readFd)
|
||||
}
|
||||
|
||||
// TestPoll_NewPoll_DoesNotCloseFdOnFailure pins the ownership rule: when
|
||||
// newPoll fails, it must leave fd open so the caller (pollQueueSet.Add's
|
||||
// callers in tun_linux.go) is the sole closer. If newPoll also closed fd,
|
||||
// the poll path would double-close on Add error. We force the failure with
|
||||
// an O_PATH descriptor: fcntl(F_SETFL) — which SetNonblock performs — is not
|
||||
// permitted on O_PATH fds and fails with EBADF, while the fd itself stays
|
||||
// open so we can observe that newPoll left it alone.
|
||||
func TestPoll_NewPoll_DoesNotCloseFdOnFailure(t *testing.T) {
|
||||
fd, err := unix.Open("/", unix.O_PATH|unix.O_CLOEXEC, 0)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = unix.Close(fd) })
|
||||
|
||||
p, err := newPoll(fd, 1)
|
||||
require.Error(t, err, "SetNonblock on an O_PATH fd should fail")
|
||||
require.Nil(t, p)
|
||||
|
||||
// If newPoll had closed fd, F_GETFD would report it closed. It staying
|
||||
// open proves newPoll left the fd for the caller to close exactly once.
|
||||
require.True(t, fdOpen(t, fd), "newPoll must not close fd on failure; caller is the sole closer")
|
||||
}
|
||||
|
||||
func TestPoll_Close_Idempotent(t *testing.T) {
|
||||
tf, err := newPoll(newReadPipe(t), 1)
|
||||
require.NoError(t, err)
|
||||
if err := tf.Close(); err != nil {
|
||||
t.Fatalf("first Close: %v", err)
|
||||
}
|
||||
if err := tf.Close(); err != nil {
|
||||
t.Fatalf("second Close should be a no-op, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// fdOpen reports whether fd currently refers to an open file description.
|
||||
// A closed (or never-allocated) fd makes F_GETFD fail with EBADF.
|
||||
func fdOpen(t *testing.T, fd int) bool {
|
||||
t.Helper()
|
||||
_, err := unix.FcntlInt(uintptr(fd), unix.F_GETFD, 0)
|
||||
if err == nil {
|
||||
return true
|
||||
}
|
||||
if errors.Is(err, unix.EBADF) {
|
||||
return false
|
||||
}
|
||||
t.Fatalf("unexpected fcntl(F_GETFD) error on fd %d: %v", fd, err)
|
||||
return false
|
||||
}
|
||||
|
||||
// TestPollQueueSet_Close_ClosesShutdownFd is the regression test for the
|
||||
// leaked shutdown eventfd: the container that owns shutdownFd must close it in
|
||||
// Close, and a second Close must be a safe no-op.
|
||||
func TestPollQueueSet_Close_ClosesShutdownFd(t *testing.T) {
|
||||
qs, err := NewPollQueueSet()
|
||||
require.NoError(t, err)
|
||||
c, ok := qs.(*pollQueueSet)
|
||||
require.True(t, ok)
|
||||
require.NoError(t, qs.Add(newReadPipe(t)))
|
||||
|
||||
shutdownFd := c.shutdownFd
|
||||
require.True(t, fdOpen(t, shutdownFd), "shutdown eventfd should be open before Close")
|
||||
|
||||
require.NoError(t, qs.Close())
|
||||
require.False(t, fdOpen(t, shutdownFd), "shutdown eventfd should be closed after Close")
|
||||
|
||||
// Second Close must not touch fds (shutdownFd is now -1) and must return nil.
|
||||
require.NoError(t, qs.Close())
|
||||
}
|
||||
|
||||
// TestOffloadQueueSet_Close_ClosesShutdownFd mirrors the poll regression test
|
||||
// for the GSO/offload queueset.
|
||||
func TestOffloadQueueSet_Close_ClosesShutdownFd(t *testing.T) {
|
||||
qs, err := NewOffloadQueueSet(false, slog.New(slog.DiscardHandler))
|
||||
require.NoError(t, err)
|
||||
c, ok := qs.(*offloadQueueSet)
|
||||
require.True(t, ok)
|
||||
require.NoError(t, qs.Add(newReadPipe(t)))
|
||||
|
||||
shutdownFd := c.shutdownFd
|
||||
require.True(t, fdOpen(t, shutdownFd), "shutdown eventfd should be open before Close")
|
||||
|
||||
require.NoError(t, qs.Close())
|
||||
require.False(t, fdOpen(t, shutdownFd), "shutdown eventfd should be closed after Close")
|
||||
|
||||
// Second Close must not touch fds (shutdownFd is now -1) and must return nil.
|
||||
require.NoError(t, qs.Close())
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio/virtio"
|
||||
)
|
||||
|
||||
// protoFromGSOType maps a virtio_net_hdr gsoType to the GSOProto value the
|
||||
// segment-time helpers use. Returns an error for GSO_NONE or any unknown
|
||||
// value. The caller should only invoke this on a confirmed superpacket.
|
||||
func protoFromGSOType(t uint8) (GSOProto, error) {
|
||||
switch t &^ unix.VIRTIO_NET_HDR_GSO_ECN {
|
||||
case unix.VIRTIO_NET_HDR_GSO_TCPV4, unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
||||
return GSOProtoTCP, nil
|
||||
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
|
||||
return GSOProtoUDP, nil
|
||||
default:
|
||||
return 0, fmt.Errorf("unsupported virtio gso type: %d", t)
|
||||
}
|
||||
}
|
||||
|
||||
// gsoTypeFromProto is the reverse of protoFromGSOType
|
||||
func gsoTypeFromProto(proto GSOProto, ipVer uint8) uint8 {
|
||||
switch {
|
||||
case proto == GSOProtoUDP && (ipVer == 4 || ipVer == 6):
|
||||
return unix.VIRTIO_NET_HDR_GSO_UDP_L4
|
||||
case ipVer == 6:
|
||||
return unix.VIRTIO_NET_HDR_GSO_TCPV6
|
||||
case ipVer == 4:
|
||||
return unix.VIRTIO_NET_HDR_GSO_TCPV4
|
||||
default:
|
||||
return unix.VIRTIO_NET_HDR_GSO_NONE
|
||||
}
|
||||
}
|
||||
|
||||
// SegmentSuperpacket invokes fn once per segment of pkt.
|
||||
// For non-GSO pkts fn is called once with pkt.Bytes.
|
||||
// For GSO/USO superpackets, fn is called once per segment with a slice of pkt.Bytes holding that segment's plaintext
|
||||
// (a freshly-patched L3+L4 header sliced in front of the original payload chunk).
|
||||
// This slicing is destructive: pkt is consumed by this call.
|
||||
// Aborts and returns the first error from fn or from per-segment construction.
|
||||
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
||||
if !pkt.GSO.IsSuperpacket() {
|
||||
return fn(pkt.Bytes)
|
||||
}
|
||||
switch pkt.GSO.Proto {
|
||||
case GSOProtoTCP:
|
||||
return virtio.SegmentTCP(pkt.Bytes, pkt.GSO.HdrLen, pkt.GSO.CsumStart, pkt.GSO.Size, fn)
|
||||
case GSOProtoUDP:
|
||||
return virtio.SegmentUDP(pkt.Bytes, pkt.GSO.HdrLen, pkt.GSO.CsumStart, pkt.GSO.Size, fn)
|
||||
default:
|
||||
return fmt.Errorf("unsupported gso proto: %d", pkt.GSO.Proto)
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,75 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package virtio
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// Size is the on-wire length of struct virtio_net_hdr the kernel
|
||||
// prepends/expects on a TUN opened with IFF_VNET_HDR (TUNSETVNETHDRSZ
|
||||
// not set).
|
||||
const Size = 10
|
||||
|
||||
// Hdr is the Go view of the legacy virtio_net_hdr.
|
||||
type Hdr struct {
|
||||
Flags uint8
|
||||
gsoType uint8 //private to avoid mistakes wrt the 0x80 VIRTIO_NET_HDR_GSO_ECN flag, ORed with the other "GSO types"
|
||||
HdrLen uint16
|
||||
GSOSize uint16
|
||||
CsumStart uint16
|
||||
CsumOffset uint16
|
||||
}
|
||||
|
||||
func NewHeader(flags, gsoType uint8, hdrLen, gsoSize, csumStart, csumOffset uint16) Hdr {
|
||||
return Hdr{
|
||||
Flags: flags,
|
||||
gsoType: gsoType,
|
||||
HdrLen: hdrLen,
|
||||
GSOSize: gsoSize,
|
||||
CsumStart: csumStart,
|
||||
CsumOffset: csumOffset,
|
||||
}
|
||||
}
|
||||
|
||||
// Decode reads a virtio_net_hdr in host byte order (TUN default; we never
|
||||
// call TUNSETVNETLE so the kernel matches our endianness).
|
||||
func (h *Hdr) Decode(b []byte) {
|
||||
h.Flags = b[0]
|
||||
h.gsoType = b[1]
|
||||
h.HdrLen = binary.NativeEndian.Uint16(b[2:4])
|
||||
h.GSOSize = binary.NativeEndian.Uint16(b[4:6])
|
||||
h.CsumStart = binary.NativeEndian.Uint16(b[6:8])
|
||||
h.CsumOffset = binary.NativeEndian.Uint16(b[8:10])
|
||||
}
|
||||
|
||||
func EncodeHeader(b []byte, flags, gsoType uint8, hdrLen, gsoSize, csumStart, csumOffset uint16) {
|
||||
b[0] = flags
|
||||
b[1] = gsoType
|
||||
binary.NativeEndian.PutUint16(b[2:4], hdrLen)
|
||||
binary.NativeEndian.PutUint16(b[4:6], gsoSize)
|
||||
binary.NativeEndian.PutUint16(b[6:8], csumStart)
|
||||
binary.NativeEndian.PutUint16(b[8:10], csumOffset)
|
||||
}
|
||||
|
||||
// Encode is the inverse of Decode: writes the virtio_net_hdr fields into b
|
||||
// (must be at least Size bytes). Used to emit a TSO superpacket on egress.
|
||||
func (h *Hdr) Encode(b []byte) {
|
||||
EncodeHeader(b, h.Flags, h.gsoType, h.HdrLen, h.GSOSize, h.CsumStart, h.CsumOffset)
|
||||
}
|
||||
|
||||
// GSOType returns gsoType with the ECN-flag masked out
|
||||
func (h *Hdr) GSOType() uint8 {
|
||||
return h.gsoType &^ unix.VIRTIO_NET_HDR_GSO_ECN
|
||||
}
|
||||
|
||||
func (h *Hdr) HasECNFlag() bool {
|
||||
return h.gsoType&unix.VIRTIO_NET_HDR_GSO_ECN != 0
|
||||
}
|
||||
|
||||
func (h *Hdr) SetGSOType(x uint8) {
|
||||
h.gsoType = x
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
//go:build !linux || android
|
||||
|
||||
package virtio
|
||||
@@ -0,0 +1,441 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
// Package virtio implements the pure validation, header-correction, and
|
||||
// per-segment slicing logic for kernel-supplied TSO/USO superpackets on
|
||||
// IFF_VNET_HDR TUN devices. It is FD-free and depends only on the byte
|
||||
// layout of the virtio_net_hdr and the IP/TCP/UDP headers it describes,
|
||||
// so it can be unit-tested in isolation from the tio Queue runtime.
|
||||
package virtio
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/checksum"
|
||||
)
|
||||
|
||||
// Protocol header size bounds used to validate / cap kernel-supplied offsets.
|
||||
const (
|
||||
ipv4HeaderMinLen = 20 // IHL=5, no options
|
||||
ipv4HeaderMaxLen = 60 // IHL=15, max options
|
||||
ipv6FixedLen = 40 // IPv6 base header; extensions would extend this
|
||||
tcpHeaderMinLen = 20 // data-offset=5, no options
|
||||
tcpHeaderMaxLen = 60 // data-offset=15, max options
|
||||
)
|
||||
|
||||
// maxSegHdrLen bounds the L3+L4 header we snapshot before stamping each segment.
|
||||
// The largest header the segmenter supports is IPv4 (max IHL 60) plus TCP (max data-offset 60) = 120 bytes
|
||||
const maxSegHdrLen = ipv4HeaderMaxLen + tcpHeaderMaxLen // 120
|
||||
|
||||
// Byte offsets inside an IPv4 header.
|
||||
const (
|
||||
ipv4TotalLenOff = 2
|
||||
ipv4IDOff = 4
|
||||
ipv4ChecksumOff = 10
|
||||
ipv4SrcOff = 12
|
||||
ipv4AddrsEnd = 20 // end of dst address (ipv4SrcOff + 2*4)
|
||||
)
|
||||
|
||||
// Byte offsets inside an IPv6 header.
|
||||
const (
|
||||
ipv6PayloadLenOff = 4
|
||||
ipv6SrcOff = 8
|
||||
ipv6AddrsEnd = 40 // end of dst address (ipv6SrcOff + 2*16)
|
||||
)
|
||||
|
||||
// Byte offsets inside a TCP header (relative to its start, i.e. csumStart).
|
||||
const (
|
||||
tcpSeqOff = 4
|
||||
tcpDataOffOff = 12 // upper nibble is header len in 32-bit words
|
||||
tcpFlagsOff = 13
|
||||
tcpChecksumOff = 16
|
||||
)
|
||||
|
||||
// UDP header is fixed at 8 bytes: {sport, dport, length, checksum}.
|
||||
const (
|
||||
udpHeaderLen = 8
|
||||
udpLengthOff = 4
|
||||
udpChecksumOff = 6
|
||||
)
|
||||
|
||||
var errPacketTooShort = errors.New("packet too short")
|
||||
|
||||
// tcpFinPshMask is cleared on every segment except the last of a TSO burst.
|
||||
const tcpFinPshMask = 0x09 // FIN(0x01) | PSH(0x08)
|
||||
|
||||
// tcpCwrFlag is cleared on every segment except the first.
|
||||
// Per RFC 3168 §6.1.2 the CWR bit signals a one-shot transition (the sender just halved its window)
|
||||
// and must appear on the first segment of a TSO burst only.
|
||||
const tcpCwrFlag = 0x80
|
||||
|
||||
// CheckValid rejects packets whose virtio_net_hdr/IP combination would
|
||||
// cause a downstream miscompute. The TUN should never emit RSC_INFO and
|
||||
// the GSO type must agree with the IP version nibble.
|
||||
func CheckValid(pkt []byte, hdr Hdr) error {
|
||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_RSC_INFO != 0 {
|
||||
return fmt.Errorf("virtio RSC_INFO flag not supported on TUN reads")
|
||||
}
|
||||
if len(pkt) < ipv4HeaderMinLen {
|
||||
return errPacketTooShort
|
||||
}
|
||||
ipVersion := pkt[0] >> 4
|
||||
if ipVersion == 6 && len(pkt) < ipv6FixedLen {
|
||||
return errPacketTooShort
|
||||
}
|
||||
|
||||
gsoType := hdr.GSOType()
|
||||
if gsoType != unix.VIRTIO_NET_HDR_GSO_NONE && hdr.GSOSize == 0 {
|
||||
// A GSO type with no segment size would dodge IsSuperpacket() downstream and
|
||||
// travel as a plain jumbo datagram with an unfinished checksum.
|
||||
return fmt.Errorf("virtio GSO type %#x with zero gso_size", hdr.gsoType)
|
||||
}
|
||||
if hdr.HasECNFlag() && !(gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV4 || gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV6) {
|
||||
return fmt.Errorf("virtio GSO_ECN qualifier on non-TCP GSO type %#x", hdr.gsoType)
|
||||
}
|
||||
switch gsoType {
|
||||
case unix.VIRTIO_NET_HDR_GSO_TCPV4:
|
||||
if ipVersion != 4 {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
}
|
||||
case unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
||||
if ipVersion != 6 {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
}
|
||||
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
|
||||
// USO carries either v4 or v6; the leading nibble disambiguates.
|
||||
if !(ipVersion == 4 || ipVersion == 6) {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
}
|
||||
default:
|
||||
if !(ipVersion == 6 || ipVersion == 4) {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// CorrectHdrLen rewrites hdr.HdrLen based on the actual transport header length read out of pkt.
|
||||
// The kernel's hdr.HdrLen on the FORWARD path can be the length of the entire first packet, so we don't trust it.
|
||||
func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
|
||||
// Thank you wireguard-go for documenting these edge-cases
|
||||
// Don't trust hdr.hdrLen from the kernel as it can be equal to the length
|
||||
// of the entire first packet when the kernel is handling it as part of a FORWARD path.
|
||||
// Instead, parse the transport header length and add it onto csumStart, which is synonymous for IP header length.
|
||||
|
||||
if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
|
||||
hdr.HdrLen = hdr.CsumStart + 8
|
||||
} else {
|
||||
if len(pkt) <= int(hdr.CsumStart+tcpDataOffOff) {
|
||||
return errors.New("packet is too short")
|
||||
}
|
||||
|
||||
tcpHLen := uint16(pkt[hdr.CsumStart+tcpDataOffOff] >> 4 * 4)
|
||||
if tcpHLen < tcpHeaderMinLen || tcpHLen > tcpHeaderMaxLen {
|
||||
return fmt.Errorf("tcp header len is invalid: %d", tcpHLen)
|
||||
}
|
||||
hdr.HdrLen = hdr.CsumStart + tcpHLen
|
||||
}
|
||||
|
||||
if len(pkt) < int(hdr.HdrLen) {
|
||||
return fmt.Errorf("length of packet (%d) < virtioNetHdr.HdrLen (%d)", len(pkt), hdr.HdrLen)
|
||||
}
|
||||
|
||||
if hdr.HdrLen < hdr.CsumStart {
|
||||
return fmt.Errorf("virtioNetHdr.HdrLen (%d) < virtioNetHdr.CsumStart (%d)", hdr.HdrLen, hdr.CsumStart)
|
||||
}
|
||||
cSumAt := int(hdr.CsumStart + hdr.CsumOffset)
|
||||
if cSumAt+1 >= len(pkt) {
|
||||
return fmt.Errorf("end of checksum offset (%d) exceeds packet length (%d)", cSumAt+1, len(pkt))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// segCount returns how many segments a payload of payLen bytes splits into at gsoSize,
|
||||
// with a floor of one so a header-only superpacket still yields a single segment.
|
||||
func segCount(payLen, gsoSize int) int {
|
||||
n := (payLen + gsoSize - 1) / gsoSize
|
||||
if n == 0 {
|
||||
return 1
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
// basePseudoSum folds the part of the L4 pseudo-header sum that is identical
|
||||
// for every segment: the source and destination addresses plus the protocol
|
||||
// number. The per-segment L4 length is added by the caller inside the loop.
|
||||
func basePseudoSum(pkt []byte, isV4 bool, proto uint32) uint32 {
|
||||
if isV4 {
|
||||
return uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0)) + proto
|
||||
}
|
||||
return uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0)) + proto
|
||||
}
|
||||
|
||||
// baseIPv4HdrSum folds the IPv4 header checksum over the fields that stay constant across segments.
|
||||
// csumStart is the L3 header length, which bounds a valid IHL.
|
||||
func baseIPv4HdrSum(pkt []byte, csumStart int) (uint32, error) {
|
||||
ihl := int(pkt[0]&0x0f) * 4
|
||||
if ihl < ipv4HeaderMinLen || ihl > csumStart {
|
||||
return 0, fmt.Errorf("bad IPv4 IHL: %d", ihl)
|
||||
}
|
||||
// total_len, the ID, and the checksum field itself are excluded: all three are rewritten per segment.
|
||||
sum := uint32(checksum.Checksum(pkt[:ihl], 0))
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4TotalLenOff : ipv4TotalLenOff+2]))
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4ChecksumOff : ipv4ChecksumOff+2]))
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2]))
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
return sum, nil
|
||||
}
|
||||
|
||||
// baseTCPHdrSum folds the TCP header checksum over everything the segment loop does not rewrite
|
||||
func baseTCPHdrSum(pkt []byte, csumStart, headerLen int) uint32 {
|
||||
seq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
|
||||
flags := uint16(pkt[csumStart+tcpFlagsOff])
|
||||
|
||||
sum := uint32(checksum.Checksum(pkt[csumStart:headerLen], 0))
|
||||
sum += uint32(^uint16(seq >> 16))
|
||||
sum += uint32(^uint16(seq))
|
||||
sum += uint32(^flags)
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[csumStart+tcpChecksumOff : csumStart+tcpChecksumOff+2]))
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
return sum
|
||||
}
|
||||
|
||||
// SegmentTCP walks a TSO superpacket pkt, yielding each segment as a slice into pkt.
|
||||
// Per-segment plaintext is laid out by stamping a copy of the original L3+L4 header into pkt at offset i*gsoSize,
|
||||
// where it sits immediately before that segment's payload chunk in the original buffer.
|
||||
// pkt is consumed by this call and must not be inspected by the caller after the final yield.
|
||||
func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
|
||||
if gsoSizeU == 0 {
|
||||
return fmt.Errorf("gso_size is zero")
|
||||
}
|
||||
if csumStartU == 0 {
|
||||
return fmt.Errorf("csum_start is zero")
|
||||
}
|
||||
|
||||
headerLen := int(hdrLenU)
|
||||
csumStart := int(csumStartU)
|
||||
if headerLen > maxSegHdrLen {
|
||||
return fmt.Errorf("header len %d exceeds max %d", headerLen, maxSegHdrLen)
|
||||
}
|
||||
isV4 := pkt[0]>>4 == 4
|
||||
|
||||
tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4
|
||||
payLen := len(pkt) - headerLen
|
||||
gsoSize := int(gsoSizeU)
|
||||
numSeg := segCount(payLen, gsoSize)
|
||||
|
||||
origSeq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
|
||||
origFlags := pkt[csumStart+tcpFlagsOff]
|
||||
|
||||
baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_TCP)
|
||||
baseTcpHdrSum := baseTCPHdrSum(pkt, csumStart, headerLen)
|
||||
|
||||
var origIPID uint16
|
||||
var baseIPHdrSum uint32
|
||||
if isV4 {
|
||||
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
|
||||
var err error
|
||||
// TSO bumps the ID per segment, so it stays out of the base sum.
|
||||
baseIPHdrSum, err = baseIPv4HdrSum(pkt, csumStart)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// Snapshot the pristine L3+L4 header once. '
|
||||
// Every segment's header is stamped from this copy, so overlapping stamps (gsoSize < headerLen) can never corrupt the source.
|
||||
var savedHdr [maxSegHdrLen]byte
|
||||
copy(savedHdr[:headerLen], pkt[:headerLen])
|
||||
|
||||
for i := 0; i < numSeg; i++ {
|
||||
segStart := i * gsoSize
|
||||
segEnd := segStart + gsoSize
|
||||
if segEnd > payLen {
|
||||
segEnd = payLen
|
||||
}
|
||||
segPayLen := segEnd - segStart
|
||||
segLen := headerLen + segPayLen
|
||||
headerOff := i * gsoSize
|
||||
|
||||
// Stamp the header into place immediately before this segment's payload, sourced from the snapshot.
|
||||
// The per-segment patches below overwrite the variable fields. (seq/flags/cksum/totalLen/id)
|
||||
if i > 0 {
|
||||
// Iter 0's header is already at pkt[:headerLen] (identical to savedHdr), so only i >= 1 needs the stamp
|
||||
copy(pkt[headerOff:headerOff+headerLen], savedHdr[:headerLen])
|
||||
}
|
||||
seg := pkt[headerOff : headerOff+segLen]
|
||||
|
||||
segSeq := origSeq + uint32(segStart)
|
||||
segFlags := origFlags
|
||||
if i != 0 {
|
||||
segFlags &^= tcpCwrFlag
|
||||
}
|
||||
if i != numSeg-1 {
|
||||
segFlags &^= tcpFinPshMask
|
||||
}
|
||||
totalLen := segLen
|
||||
|
||||
if isV4 {
|
||||
segID := origIPID + uint16(i)
|
||||
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
|
||||
binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID)
|
||||
ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID)
|
||||
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
|
||||
} else {
|
||||
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
|
||||
}
|
||||
|
||||
binary.BigEndian.PutUint32(seg[csumStart+tcpSeqOff:csumStart+tcpSeqOff+4], segSeq)
|
||||
seg[csumStart+tcpFlagsOff] = segFlags
|
||||
|
||||
tcpLen := tcpHdrLen + segPayLen
|
||||
// Payload bytes still live at their original offset in pkt.
|
||||
// The header slide above only writes into pkt[i*GSOSize : i*GSOSize+header], which is the tail of seg_{i-1}'s payload (already consumed)
|
||||
// and never overlaps seg_i's own payload at pkt[header+i*GSOSize : header+(i+1)*GSOSize].
|
||||
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
|
||||
wide := uint64(baseTcpHdrSum) + uint64(paySum) + uint64(baseProtoSum)
|
||||
wide += uint64(segSeq) + uint64(segFlags) + uint64(tcpLen)
|
||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
||||
binary.BigEndian.PutUint16(seg[csumStart+tcpChecksumOff:csumStart+tcpChecksumOff+2], foldComplement(uint32(wide)))
|
||||
|
||||
if err := yield(seg); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// SegmentUDP walks a USO superpacket, stamping a per-segment-patched copy of the original L3+L4 header
|
||||
// into pkt at offset i*GSOSize and yielding pkt[i*GSOSize:i*GSOSize+segLen] to the caller.
|
||||
// Per-segment patches are total_len + IPv4 csum (or IPv6 payload_len) plus the UDP length and checksum.
|
||||
// pkt is consumed destructively.
|
||||
func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
|
||||
if gsoSizeU == 0 {
|
||||
return fmt.Errorf("gso_size is zero")
|
||||
}
|
||||
if csumStartU == 0 {
|
||||
return fmt.Errorf("csum_start is zero")
|
||||
}
|
||||
|
||||
isV4 := pkt[0]>>4 == 4
|
||||
headerLen := int(hdrLenU)
|
||||
csumStart := int(csumStartU)
|
||||
if headerLen > maxSegHdrLen {
|
||||
return fmt.Errorf("header len %d exceeds max %d", headerLen, maxSegHdrLen)
|
||||
}
|
||||
if headerLen-csumStart != udpHeaderLen {
|
||||
return fmt.Errorf("udp header len mismatch: %d", headerLen-csumStart)
|
||||
}
|
||||
|
||||
payLen := len(pkt) - headerLen
|
||||
gsoSize := int(gsoSizeU)
|
||||
numSeg := segCount(payLen, gsoSize)
|
||||
|
||||
baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_UDP)
|
||||
|
||||
var origIPID uint16
|
||||
var baseIPHdrSum uint32
|
||||
if isV4 {
|
||||
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
|
||||
var err error
|
||||
// Software UDP GSO bumps the ID per segment just like TSO
|
||||
// (inet_gso_segment's fixed-ID case is TCP-only), so it stays out of the base sum.
|
||||
baseIPHdrSum, err = baseIPv4HdrSum(pkt, csumStart)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// Snapshot the pristine L3+L4 header once and stamp every segment from it
|
||||
var savedHdr [maxSegHdrLen]byte
|
||||
copy(savedHdr[:headerLen], pkt[:headerLen])
|
||||
|
||||
for i := 0; i < numSeg; i++ {
|
||||
segStart := i * gsoSize
|
||||
segEnd := segStart + gsoSize
|
||||
if segEnd > payLen {
|
||||
segEnd = payLen
|
||||
}
|
||||
segPayLen := segEnd - segStart
|
||||
segLen := headerLen + segPayLen
|
||||
headerOff := i * gsoSize
|
||||
|
||||
if i > 0 {
|
||||
copy(pkt[headerOff:headerOff+headerLen], savedHdr[:headerLen])
|
||||
}
|
||||
seg := pkt[headerOff : headerOff+segLen]
|
||||
|
||||
totalLen := segLen
|
||||
udpLen := udpHeaderLen + segPayLen
|
||||
|
||||
if isV4 {
|
||||
segID := origIPID + uint16(i)
|
||||
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
|
||||
binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID)
|
||||
ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID)
|
||||
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
|
||||
} else {
|
||||
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
|
||||
}
|
||||
|
||||
binary.BigEndian.PutUint16(seg[csumStart+udpLengthOff:csumStart+udpLengthOff+2], uint16(udpLen))
|
||||
|
||||
// Sum the UDP header (length just written, checksum zeroed) together with
|
||||
// this segment's payload in one pass, seeded with the pseudo-header sum.
|
||||
seg[csumStart+udpChecksumOff], seg[csumStart+udpChecksumOff+1] = 0, 0
|
||||
pseudo := baseProtoSum + uint32(udpLen)
|
||||
pseudo = (pseudo & 0xffff) + (pseudo >> 16)
|
||||
pseudo = (pseudo & 0xffff) + (pseudo >> 16)
|
||||
csum := ^checksum.Checksum(seg[csumStart:], uint16(pseudo))
|
||||
if csum == 0 {
|
||||
csum = 0xffff
|
||||
}
|
||||
binary.BigEndian.PutUint16(seg[csumStart+udpChecksumOff:csumStart+udpChecksumOff+2], csum)
|
||||
|
||||
if err := yield(seg); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// FinishChecksum computes the L4 checksum for a non-GSO packet that the kernel handed us with NEEDS_CSUM set.
|
||||
// CsumStart / CsumOffset point at the 16-bit checksum field.
|
||||
// We zero it, fold a full sum from the partial one that the kernel provided, and store the result.
|
||||
func FinishChecksum(seg []byte, hdr Hdr) error {
|
||||
cs := int(hdr.CsumStart)
|
||||
co := int(hdr.CsumOffset)
|
||||
if cs+co+2 > len(seg) {
|
||||
return fmt.Errorf("csum offsets out of range: start=%d offset=%d len=%d", cs, co, len(seg))
|
||||
}
|
||||
// The kernel stores a partial pseudo-header sum at [cs+co:]; sum over the
|
||||
// L4 region starting at cs, folding the prior partial in as the seed.
|
||||
partial := binary.BigEndian.Uint16(seg[cs+co : cs+co+2])
|
||||
seg[cs+co] = 0
|
||||
seg[cs+co+1] = 0
|
||||
csum := ^checksum.Checksum(seg[cs:], partial)
|
||||
// RFC 768: UDP transmits a computed zero as all ones, since all-zero is the reserved "no checksum" value.
|
||||
if co == udpChecksumOff && csum == 0 {
|
||||
csum = 0xffff
|
||||
}
|
||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], csum)
|
||||
return nil
|
||||
}
|
||||
|
||||
// foldComplement folds a 32-bit one's-complement partial sum to 16 bits and
|
||||
// complements it, yielding the on-wire Internet checksum value.
|
||||
func foldComplement(sum uint32) uint16 {
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
return ^uint16(sum)
|
||||
}
|
||||
@@ -0,0 +1,602 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package virtio
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"testing"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/checksum"
|
||||
)
|
||||
|
||||
// verifyChecksum confirms that the one's-complement sum across b, seeded with
|
||||
// a folded pseudo-header sum, equals all-ones (a valid on-wire checksum).
|
||||
// A corrupted header stamped into a segment makes this fail even when the
|
||||
// checksum field itself was computed from the (pristine) base sums, because
|
||||
// the bytes the receiver would sum no longer match what was checksummed.
|
||||
func verifyChecksum(b []byte, pseudo uint16) bool {
|
||||
return checksum.Checksum(b, pseudo) == 0xffff
|
||||
}
|
||||
|
||||
// pseudoHeaderIPv4 folds the TCP/UDP pseudo-header sum from a segment's own
|
||||
// address and length fields, used to independently verify its L4 checksum.
|
||||
func pseudoHeaderIPv4(src, dst []byte, proto byte, l4Len int) uint16 {
|
||||
s := uint32(checksum.Checksum(src, 0)) + uint32(checksum.Checksum(dst, 0))
|
||||
s += uint32(proto) + uint32(l4Len)
|
||||
s = (s & 0xffff) + (s >> 16)
|
||||
s = (s & 0xffff) + (s >> 16)
|
||||
return uint16(s)
|
||||
}
|
||||
|
||||
// buildTCPv4Super constructs a synthetic IPv4/TCP TSO superpacket with a
|
||||
// payload of payLen bytes and returns it alongside the header fields the
|
||||
// segmenter needs. The header is a fixed 40 bytes (20 IPv4 + 20 TCP).
|
||||
func buildTCPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
|
||||
const ipLen = 20
|
||||
const tcpLen = 20
|
||||
pkt = make([]byte, ipLen+tcpLen+payLen)
|
||||
|
||||
// IPv4 header.
|
||||
pkt[0] = 0x45 // version 4, IHL 5
|
||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+tcpLen+payLen))
|
||||
binary.BigEndian.PutUint16(pkt[4:6], 0x4242) // ID
|
||||
pkt[8] = 64 // TTL
|
||||
pkt[9] = unix.IPPROTO_TCP
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1}) // src
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2}) // dst
|
||||
|
||||
// TCP header.
|
||||
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
|
||||
binary.BigEndian.PutUint16(pkt[22:24], 80) // dport
|
||||
binary.BigEndian.PutUint32(pkt[24:28], 10000) // seq
|
||||
binary.BigEndian.PutUint32(pkt[28:32], 20000) // ack
|
||||
pkt[32] = 0x50 // data offset 5 words
|
||||
pkt[33] = 0x18 // ACK | PSH
|
||||
binary.BigEndian.PutUint16(pkt[34:36], 65535) // window
|
||||
|
||||
for i := 0; i < payLen; i++ {
|
||||
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
|
||||
}
|
||||
return pkt, ipLen + tcpLen, ipLen
|
||||
}
|
||||
|
||||
// buildUDPv4Super constructs a synthetic IPv4/UDP USO superpacket with a
|
||||
// payload of payLen bytes. Header is a fixed 28 bytes (20 IPv4 + 8 UDP).
|
||||
func buildUDPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
|
||||
const ipLen = 20
|
||||
const udpLen = 8
|
||||
pkt = make([]byte, ipLen+udpLen+payLen)
|
||||
|
||||
pkt[0] = 0x45
|
||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+udpLen+payLen))
|
||||
binary.BigEndian.PutUint16(pkt[4:6], 0x4242)
|
||||
pkt[8] = 64
|
||||
pkt[9] = unix.IPPROTO_UDP
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||
|
||||
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
|
||||
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
|
||||
|
||||
for i := 0; i < payLen; i++ {
|
||||
pkt[ipLen+udpLen+i] = byte(i & 0xff)
|
||||
}
|
||||
return pkt, ipLen + udpLen, ipLen
|
||||
}
|
||||
|
||||
// collectTCP segments a fresh copy of pkt and returns each segment as an
|
||||
// independent slice so assertions can run after segmentation completes.
|
||||
func collectTCP(t *testing.T, pkt []byte, hdrLen, csumStart, gsoSize uint16) [][]byte {
|
||||
t.Helper()
|
||||
work := append([]byte(nil), pkt...)
|
||||
var out [][]byte
|
||||
err := SegmentTCP(work, hdrLen, csumStart, gsoSize, func(seg []byte) error {
|
||||
out = append(out, append([]byte(nil), seg...))
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SegmentTCP: %v", err)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func collectUDP(t *testing.T, pkt []byte, hdrLen, csumStart, gsoSize uint16) [][]byte {
|
||||
t.Helper()
|
||||
work := append([]byte(nil), pkt...)
|
||||
var out [][]byte
|
||||
err := SegmentUDP(work, hdrLen, csumStart, gsoSize, func(seg []byte) error {
|
||||
out = append(out, append([]byte(nil), seg...))
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SegmentUDP: %v", err)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// TestSegmentTCPHeaderNotCorrupted is the regression test for the in-place
|
||||
// header-slide bug: when gsoSize < headerLen the old code stamped each
|
||||
// segment's header from pkt[:headerLen], which had already been overwritten
|
||||
// by the previous segment's overlapping stamp, so segments 2..n carried a
|
||||
// corrupted header (garbage src/dst/ports/seq). Every segment must instead
|
||||
// carry the ORIGINAL constant header fields with correct per-segment seq.
|
||||
func TestSegmentTCPHeaderNotCorrupted(t *testing.T) {
|
||||
const origSeq = 10000
|
||||
cases := []struct {
|
||||
name string
|
||||
payLen int
|
||||
gsoSize uint16
|
||||
}{
|
||||
// gsoSize (8) < headerLen (40): the bug's trigger. Even split.
|
||||
{"small-gso-even", 40, 8},
|
||||
// gsoSize (8) < headerLen (40) with a short final segment.
|
||||
{"small-gso-odd-tail", 44, 8},
|
||||
// gsoSize (100) >= headerLen (40): the normal path, must still work.
|
||||
{"normal-gso", 250, 100},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
pkt, hdrLen, csumStart := buildTCPv4Super(tc.payLen)
|
||||
gso := int(tc.gsoSize)
|
||||
wantSeg := (tc.payLen + gso - 1) / gso
|
||||
segs := collectTCP(t, pkt, hdrLen, csumStart, tc.gsoSize)
|
||||
if len(segs) != wantSeg {
|
||||
t.Fatalf("got %d segments, want %d", len(segs), wantSeg)
|
||||
}
|
||||
|
||||
off := 0
|
||||
for i, seg := range segs {
|
||||
// Constant header fields must be identical to the original in
|
||||
// EVERY segment. These are exactly the bytes the old code
|
||||
// corrupted in segments 2..n.
|
||||
if got := seg[0]; got != 0x45 {
|
||||
t.Errorf("seg %d: version/IHL byte=%#x want 0x45", i, got)
|
||||
}
|
||||
if seg[9] != unix.IPPROTO_TCP {
|
||||
t.Errorf("seg %d: proto=%d want %d", i, seg[9], unix.IPPROTO_TCP)
|
||||
}
|
||||
if !bytes.Equal(seg[12:16], []byte{10, 0, 0, 1}) {
|
||||
t.Errorf("seg %d: src=%v want [10 0 0 1]", i, seg[12:16])
|
||||
}
|
||||
if !bytes.Equal(seg[16:20], []byte{10, 0, 0, 2}) {
|
||||
t.Errorf("seg %d: dst=%v want [10 0 0 2]", i, seg[16:20])
|
||||
}
|
||||
if sport := binary.BigEndian.Uint16(seg[20:22]); sport != 12345 {
|
||||
t.Errorf("seg %d: sport=%d want 12345", i, sport)
|
||||
}
|
||||
if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 80 {
|
||||
t.Errorf("seg %d: dport=%d want 80", i, dport)
|
||||
}
|
||||
if ack := binary.BigEndian.Uint32(seg[28:32]); ack != 20000 {
|
||||
t.Errorf("seg %d: ack=%d want 20000", i, ack)
|
||||
}
|
||||
if seg[32] != 0x50 {
|
||||
t.Errorf("seg %d: data-offset byte=%#x want 0x50", i, seg[32])
|
||||
}
|
||||
|
||||
// Per-segment seq must advance by the payload offset.
|
||||
segStart := i * gso
|
||||
if seq := binary.BigEndian.Uint32(seg[24:28]); seq != uint32(origSeq+segStart) {
|
||||
t.Errorf("seg %d: seq=%d want %d", i, seq, origSeq+segStart)
|
||||
}
|
||||
|
||||
// Payload bytes must be the original contiguous slice.
|
||||
segPayLen := len(seg) - int(hdrLen)
|
||||
wantPay := make([]byte, segPayLen)
|
||||
for k := 0; k < segPayLen; k++ {
|
||||
wantPay[k] = byte((off + k) & 0xff)
|
||||
}
|
||||
if !bytes.Equal(seg[hdrLen:], wantPay) {
|
||||
t.Errorf("seg %d: payload mismatch", i)
|
||||
}
|
||||
off += segPayLen
|
||||
|
||||
// End-to-end: the stamped header must checksum-verify. A
|
||||
// corrupted header fails here because the written checksum was
|
||||
// derived from the pristine header.
|
||||
if !verifyChecksum(seg[:20], 0) {
|
||||
t.Errorf("seg %d: bad IPv4 header checksum", i)
|
||||
}
|
||||
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_TCP, len(seg)-20)
|
||||
if !verifyChecksum(seg[20:], psum) {
|
||||
t.Errorf("seg %d: bad TCP checksum", i)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestCorrectHdrLenChecksumBound guards the checksum-field bounds check in
|
||||
// CorrectHdrLen. The checksum field sits at CsumStart+CsumOffset, so the check
|
||||
// must be computed from CsumStart+CsumOffset — NOT CsumStart+CsumStart, a
|
||||
// regression that doubled CsumStart and thus over-tightened the bound (since
|
||||
// CsumOffset, 6 for UDP / 16 for TCP, is always < CsumStart >= 20). That bogus
|
||||
// bound spuriously rejected valid small USO superpackets in decodeRead.
|
||||
func TestCorrectHdrLenChecksumBound(t *testing.T) {
|
||||
// A valid IPv4 USO superpacket: 20B IPv4 + 8B UDP + two 6-byte segments
|
||||
// (payload 12) = 40 bytes total. CsumStart=20, CsumOffset=6, so the UDP
|
||||
// checksum field lives at bytes 26..27, comfortably inside the 40-byte
|
||||
// packet. The OLD formula computed cSumAt = CsumStart+CsumStart = 40 and
|
||||
// rejected on cSumAt+1 (41) >= len(pkt) (40); the fix (CsumStart+CsumOffset
|
||||
// = 26) accepts. This case FAILS against the CsumStart+CsumStart regression.
|
||||
t.Run("valid-small-uso-accepted", func(t *testing.T) {
|
||||
pkt, _, csumStart := buildUDPv4Super(12) // total len 40
|
||||
hdr := NewHeader(
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
|
||||
0, /*hdrLen*/
|
||||
6, /*gsoSize: two 6-byte segments*/
|
||||
csumStart, /*csumStart*/
|
||||
6, /*csumOffset*/
|
||||
)
|
||||
if err := CorrectHdrLen(pkt, &hdr); err != nil {
|
||||
t.Fatalf("CorrectHdrLen rejected a valid 40-byte USO superpacket: %v", err)
|
||||
}
|
||||
if hdr.HdrLen != csumStart+udpHeaderLen {
|
||||
t.Errorf("HdrLen = %d, want %d", hdr.HdrLen, csumStart+udpHeaderLen)
|
||||
}
|
||||
})
|
||||
|
||||
// A genuinely-too-short packet: CsumStart=20, CsumOffset=6 means the
|
||||
// checksum field would end at byte 27, but the packet is only 25 bytes
|
||||
// (CsumStart+CsumOffset+2 = 28 > 25). CorrectHdrLen must still reject it.
|
||||
t.Run("too-short-rejected", func(t *testing.T) {
|
||||
pkt := make([]byte, 25)
|
||||
pkt[0] = 0x45 // IPv4, IHL 5
|
||||
hdr := NewHeader(
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
|
||||
0, /*hdrLen*/
|
||||
6, /*gsoSize*/
|
||||
20, /*csumStart*/
|
||||
6, /*csumOffset*/
|
||||
)
|
||||
if err := CorrectHdrLen(pkt, &hdr); err == nil {
|
||||
t.Fatalf("CorrectHdrLen accepted a too-short (25-byte) packet")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestSegmentUDPHeaderNotCorrupted is the USO counterpart: SegmentUDP performs
|
||||
// the same header stamp and must be correct when gsoSize < headerLen.
|
||||
func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
payLen int
|
||||
gsoSize uint16
|
||||
}{
|
||||
{"small-gso-even", 40, 8},
|
||||
{"small-gso-odd-tail", 44, 8},
|
||||
{"normal-gso", 250, 100},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
pkt, hdrLen, csumStart := buildUDPv4Super(tc.payLen)
|
||||
gso := int(tc.gsoSize)
|
||||
wantSeg := (tc.payLen + gso - 1) / gso
|
||||
segs := collectUDP(t, pkt, hdrLen, csumStart, tc.gsoSize)
|
||||
if len(segs) != wantSeg {
|
||||
t.Fatalf("got %d segments, want %d", len(segs), wantSeg)
|
||||
}
|
||||
|
||||
off := 0
|
||||
for i, seg := range segs {
|
||||
if got := seg[0]; got != 0x45 {
|
||||
t.Errorf("seg %d: version/IHL byte=%#x want 0x45", i, got)
|
||||
}
|
||||
if seg[9] != unix.IPPROTO_UDP {
|
||||
t.Errorf("seg %d: proto=%d want %d", i, seg[9], unix.IPPROTO_UDP)
|
||||
}
|
||||
if !bytes.Equal(seg[12:16], []byte{10, 0, 0, 1}) {
|
||||
t.Errorf("seg %d: src=%v want [10 0 0 1]", i, seg[12:16])
|
||||
}
|
||||
if !bytes.Equal(seg[16:20], []byte{10, 0, 0, 2}) {
|
||||
t.Errorf("seg %d: dst=%v want [10 0 0 2]", i, seg[16:20])
|
||||
}
|
||||
if sport := binary.BigEndian.Uint16(seg[20:22]); sport != 12345 {
|
||||
t.Errorf("seg %d: sport=%d want 12345", i, sport)
|
||||
}
|
||||
if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 53 {
|
||||
t.Errorf("seg %d: dport=%d want 53", i, dport)
|
||||
}
|
||||
// Software UDP GSO bumps the IPv4 ID per segment just like TSO
|
||||
// (inet_gso_segment's fixed-ID case is TCP-only).
|
||||
if id := binary.BigEndian.Uint16(seg[4:6]); id != 0x4242+uint16(i) {
|
||||
t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242+uint16(i))
|
||||
}
|
||||
|
||||
segPayLen := len(seg) - int(hdrLen)
|
||||
if udpLen := binary.BigEndian.Uint16(seg[24:26]); udpLen != uint16(8+segPayLen) {
|
||||
t.Errorf("seg %d: udp len=%d want %d", i, udpLen, 8+segPayLen)
|
||||
}
|
||||
|
||||
wantPay := make([]byte, segPayLen)
|
||||
for k := 0; k < segPayLen; k++ {
|
||||
wantPay[k] = byte((off + k) & 0xff)
|
||||
}
|
||||
if !bytes.Equal(seg[hdrLen:], wantPay) {
|
||||
t.Errorf("seg %d: payload mismatch", i)
|
||||
}
|
||||
off += segPayLen
|
||||
|
||||
if !verifyChecksum(seg[:20], 0) {
|
||||
t.Errorf("seg %d: bad IPv4 header checksum", i)
|
||||
}
|
||||
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_UDP, len(seg)-20)
|
||||
if !verifyChecksum(seg[20:], psum) {
|
||||
t.Errorf("seg %d: bad UDP checksum", i)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// buildUDPv4Single constructs a single IPv4/UDP datagram with the checksum field pre-loaded with the folded
|
||||
// pseudo-header sum, which is how a NEEDS_CSUM (CHECKSUM_PARTIAL) packet arrives from the tun.
|
||||
func buildUDPv4Single(payload []byte) (pkt []byte, hdr Hdr) {
|
||||
const ipLen, udpLen = 20, 8
|
||||
pkt = make([]byte, ipLen+udpLen+len(payload))
|
||||
|
||||
pkt[0] = 0x45
|
||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(len(pkt)))
|
||||
pkt[8] = 64
|
||||
pkt[9] = unix.IPPROTO_UDP
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||
|
||||
binary.BigEndian.PutUint16(pkt[ipLen:ipLen+2], 12345)
|
||||
binary.BigEndian.PutUint16(pkt[ipLen+2:ipLen+4], 53)
|
||||
binary.BigEndian.PutUint16(pkt[ipLen+4:ipLen+6], uint16(udpLen+len(payload)))
|
||||
copy(pkt[ipLen+udpLen:], payload)
|
||||
|
||||
pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, udpLen+len(payload))
|
||||
binary.BigEndian.PutUint16(pkt[ipLen+udpChecksumOff:ipLen+udpChecksumOff+2], pseudo)
|
||||
|
||||
return pkt, NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, ipLen, udpChecksumOff)
|
||||
}
|
||||
|
||||
// TestFinishChecksumUDPZeroStoresAllOnes pins RFC 768: a UDP checksum that computes to zero goes on the wire as
|
||||
// 0xffff, because all-zero is the reserved "no checksum" encoding that IPv6 rejects outright.
|
||||
func TestFinishChecksumUDPZeroStoresAllOnes(t *testing.T) {
|
||||
var payload []byte
|
||||
for i := 0; i < 0x10000; i++ {
|
||||
p := []byte{byte(i >> 8), byte(i)}
|
||||
pkt, hdr := buildUDPv4Single(p)
|
||||
cs, co := int(hdr.CsumStart), int(hdr.CsumOffset)
|
||||
partial := binary.BigEndian.Uint16(pkt[cs+co : cs+co+2])
|
||||
pkt[cs+co], pkt[cs+co+1] = 0, 0
|
||||
if ^checksum.Checksum(pkt[cs:], partial) == 0 {
|
||||
payload = p
|
||||
break
|
||||
}
|
||||
}
|
||||
if payload == nil {
|
||||
t.Fatal("no 2-byte payload produced a zero checksum")
|
||||
}
|
||||
|
||||
pkt, hdr := buildUDPv4Single(payload)
|
||||
if err := FinishChecksum(pkt, hdr); err != nil {
|
||||
t.Fatalf("FinishChecksum: %v", err)
|
||||
}
|
||||
off := int(hdr.CsumStart) + int(hdr.CsumOffset)
|
||||
if got := binary.BigEndian.Uint16(pkt[off : off+2]); got != 0xffff {
|
||||
t.Fatalf("stored %#04x, want 0xffff for a UDP checksum that folds to zero", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFinishChecksumTCPZeroPreserved is the other half of the protocol split: zero is a legal TCP checksum and must
|
||||
// be stored as-is, so the UDP rewrite must not fire on a TCP csum_offset.
|
||||
func TestFinishChecksumTCPZeroPreserved(t *testing.T) {
|
||||
const cs, co = 20, tcpChecksumOff
|
||||
|
||||
// Seed the partial so the completed checksum lands on zero, the case the UDP path rewrites.
|
||||
seg := make([]byte, cs+co+2)
|
||||
for i := range seg[cs:] {
|
||||
seg[cs+i] = byte(i * 7)
|
||||
}
|
||||
var partial uint16
|
||||
for i := 0; i <= 0xffff; i++ {
|
||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], uint16(i))
|
||||
probe := append([]byte(nil), seg...)
|
||||
probe[cs+co], probe[cs+co+1] = 0, 0
|
||||
if ^checksum.Checksum(probe[cs:], uint16(i)) == 0 {
|
||||
partial = uint16(i)
|
||||
break
|
||||
}
|
||||
}
|
||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], partial)
|
||||
|
||||
hdr := NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, cs, co)
|
||||
if err := FinishChecksum(seg, hdr); err != nil {
|
||||
t.Fatalf("FinishChecksum: %v", err)
|
||||
}
|
||||
if got := binary.BigEndian.Uint16(seg[cs+co : cs+co+2]); got != 0 {
|
||||
t.Fatalf("stored %#04x, want 0x0000 (zero is a legal TCP checksum)", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFinishChecksumUDPValidates confirms the ordinary path still produces a checksum a receiver accepts.
|
||||
func TestFinishChecksumUDPValidates(t *testing.T) {
|
||||
payload := []byte("the definitive tun offloads branch")
|
||||
pkt, hdr := buildUDPv4Single(payload)
|
||||
if err := FinishChecksum(pkt, hdr); err != nil {
|
||||
t.Fatalf("FinishChecksum: %v", err)
|
||||
}
|
||||
pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, 8+len(payload))
|
||||
if !verifyChecksum(pkt[hdr.CsumStart:], pseudo) {
|
||||
t.Fatal("completed UDP checksum does not validate")
|
||||
}
|
||||
}
|
||||
|
||||
// TestCheckValidMasksGSOECN: GSO_ECN is a qualifier bit the kernel ORs
|
||||
// into gso_type for TSO superpackets with CWR set. CheckValid must
|
||||
// validate an ECN-qualified type as its base type — previously TCPV4|ECN
|
||||
// fell into the default case and skipped the IP-version agreement check.
|
||||
// The qualifier is TCP-only, so it must be rejected on UDP_L4.
|
||||
func TestCheckValidMasksGSOECN(t *testing.T) {
|
||||
v4pkt, _, _ := buildTCPv4Super(100)
|
||||
v6pkt := make([]byte, len(v4pkt))
|
||||
copy(v6pkt, v4pkt)
|
||||
v6pkt[0] = 0x60 // claim IPv6
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
pkt []byte
|
||||
gsoType uint8
|
||||
wantErr bool
|
||||
}{
|
||||
{"tcpv4-ecn-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, false},
|
||||
{"tcpv4-ecn-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
|
||||
{"tcpv6-ecn-v4-mismatch", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV6 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
|
||||
{"udp-l4-ecn-rejected", v4pkt, unix.VIRTIO_NET_HDR_GSO_UDP_L4 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
|
||||
{"tcpv4-plain-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, false},
|
||||
{"tcpv4-plain-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, true},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
err := CheckValid(tc.pkt, NewHeader(0, tc.gsoType, 0, 100, 0, 0))
|
||||
if tc.wantErr && err == nil {
|
||||
t.Errorf("CheckValid(gsoType=%#x) = nil, want error", tc.gsoType)
|
||||
}
|
||||
if !tc.wantErr && err != nil {
|
||||
t.Errorf("CheckValid(gsoType=%#x) = %v, want nil", tc.gsoType, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestCheckValidRejectsZeroGSOSize: a GSO-typed header with gso_size=0 must be
|
||||
// rejected. It would produce a Packet whose GSOInfo.IsSuperpacket() is false,
|
||||
// dodging both segmentation and FinishChecksum on its way downstream.
|
||||
func TestCheckValidRejectsZeroGSOSize(t *testing.T) {
|
||||
v4pkt, _, _ := buildTCPv4Super(100)
|
||||
if err := CheckValid(v4pkt, NewHeader(0, unix.VIRTIO_NET_HDR_GSO_TCPV4, 0, 0, 0, 0)); err == nil {
|
||||
t.Fatal("CheckValid accepted a GSO-typed header with gso_size=0")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFoldComplementMatchesReference checks the segmenter's fold-and-invert
|
||||
// against an independent RFC 1071 reference fold, hitting the carry edge
|
||||
// cases (values whose first fold produces another carry).
|
||||
func TestFoldComplementMatchesReference(t *testing.T) {
|
||||
refFold := func(s uint64) uint16 {
|
||||
for s>>16 != 0 {
|
||||
s = s&0xffff + s>>16
|
||||
}
|
||||
return uint16(s)
|
||||
}
|
||||
cases := []uint32{
|
||||
0, 1, 0xffff,
|
||||
0x10000, // single carry
|
||||
0x1fffe, // 0xffff + 0xffff: carry produces another 0xffff
|
||||
0xffff0000, // high half only
|
||||
0xfffeffff, // first fold yields another carry
|
||||
0xffffffff, // worst case
|
||||
}
|
||||
for _, c := range cases {
|
||||
if got, want := foldComplement(c), ^refFold(uint64(c)); got != want {
|
||||
t.Errorf("foldComplement(%#x) = %#x, want %#x", c, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// referenceBaseIPv4HdrSum and referenceBaseTCPHdrSum are the straightforward
|
||||
// implementations that baseIPv4HdrSum/baseTCPHdrSum replaced: copy the header
|
||||
// into scratch, zero the fields the segment loop rewrites, sum. The production
|
||||
// versions instead sum in place and subtract those fields via one's-complement
|
||||
// arithmetic, which is faster but far less obvious — particularly for the TCP
|
||||
// flags byte, which is only half of a 16-bit word. These references exist so
|
||||
// that trade is checked rather than asserted.
|
||||
func referenceBaseIPv4HdrSum(pkt []byte, ihl int) uint32 {
|
||||
var ipTmp [ipv4HeaderMaxLen]byte
|
||||
copy(ipTmp[:ihl], pkt[:ihl])
|
||||
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
|
||||
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
|
||||
ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0
|
||||
return uint32(checksum.Checksum(ipTmp[:ihl], 0))
|
||||
}
|
||||
|
||||
func referenceBaseTCPHdrSum(pkt []byte, csumStart, headerLen int) uint32 {
|
||||
tcpLen := headerLen - csumStart
|
||||
var tmp [tcpHeaderMaxLen]byte
|
||||
copy(tmp[:tcpLen], pkt[csumStart:headerLen])
|
||||
tmp[tcpSeqOff], tmp[tcpSeqOff+1], tmp[tcpSeqOff+2], tmp[tcpSeqOff+3] = 0, 0, 0, 0
|
||||
tmp[tcpFlagsOff] = 0
|
||||
tmp[tcpChecksumOff], tmp[tcpChecksumOff+1] = 0, 0
|
||||
return uint32(checksum.Checksum(tmp[:tcpLen], 0))
|
||||
}
|
||||
|
||||
// randSeed is a tiny deterministic PRNG so this test needs no imports beyond
|
||||
// what the file already has and reproduces identically on every run.
|
||||
func randByte(state *uint32) byte {
|
||||
*state = *state*1664525 + 1013904223
|
||||
return byte(*state >> 24)
|
||||
}
|
||||
|
||||
func TestBaseSumsMatchZeroingReference(t *testing.T) {
|
||||
state := uint32(12345)
|
||||
|
||||
t.Run("ipv4", func(t *testing.T) {
|
||||
for ihl := ipv4HeaderMinLen; ihl <= ipv4HeaderMaxLen; ihl += 4 {
|
||||
for iter := 0; iter < 5000; iter++ {
|
||||
pkt := make([]byte, ihl)
|
||||
for i := range pkt {
|
||||
pkt[i] = randByte(&state)
|
||||
}
|
||||
pkt[0] = byte(0x40 | (ihl / 4))
|
||||
|
||||
want := referenceBaseIPv4HdrSum(pkt, ihl)
|
||||
got, err := baseIPv4HdrSum(pkt, ihl)
|
||||
if err != nil {
|
||||
t.Fatalf("ihl=%d: %v", ihl, err)
|
||||
}
|
||||
// Compare the value that reaches the wire: the raw partial
|
||||
// sums may legally differ by one's-complement -0 vs +0.
|
||||
for _, tl := range []uint32{20, 1500, 65535} {
|
||||
for _, id := range []uint32{0, 0x4242, 0xffff} {
|
||||
if a, b := foldComplement(want+tl+id), foldComplement(got+tl+id); a != b {
|
||||
t.Fatalf("ihl=%d tl=%d id=%d: %#04x != %#04x", ihl, tl, id, a, b)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("tcp", func(t *testing.T) {
|
||||
const csumStart = 20
|
||||
for dataOff := 5; dataOff <= 15; dataOff++ {
|
||||
tcpLen := dataOff * 4
|
||||
headerLen := csumStart + tcpLen
|
||||
for iter := 0; iter < 5000; iter++ {
|
||||
pkt := make([]byte, headerLen+64)
|
||||
for i := range pkt {
|
||||
pkt[i] = randByte(&state)
|
||||
}
|
||||
pkt[0] = 0x45
|
||||
pkt[csumStart+tcpDataOffOff] = byte(dataOff << 4)
|
||||
|
||||
want := referenceBaseTCPHdrSum(pkt, csumStart, headerLen)
|
||||
got := baseTCPHdrSum(pkt, csumStart, headerLen)
|
||||
for _, seq := range []uint32{0, 1, 0x4242_4242, 0xffff_ffff} {
|
||||
for _, fl := range []uint32{0x00, 0x10, 0x18, 0x19, 0xff} {
|
||||
for _, l4 := range []uint32{20, 1460, 65535} {
|
||||
a := foldComplement(want + seq + fl + l4)
|
||||
b := foldComplement(got + seq + fl + l4)
|
||||
if a != b {
|
||||
t.Fatalf("dataOff=%d seq=%#x fl=%#x l4=%d: %#04x != %#04x",
|
||||
dataOff, seq, fl, l4, a, b)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
)
|
||||
@@ -63,7 +64,7 @@ func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
||||
return r
|
||||
}
|
||||
|
||||
func (t tun) Activate() error {
|
||||
func (t *tun) Activate() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -96,10 +97,6 @@ func (t *tun) Name() string {
|
||||
return "android"
|
||||
}
|
||||
|
||||
func (t *tun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for android")
|
||||
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||
}
|
||||
|
||||
@@ -6,7 +6,6 @@ package overlay
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"os"
|
||||
@@ -16,6 +15,7 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
netroute "golang.org/x/net/route"
|
||||
@@ -550,7 +550,9 @@ func (t *tun) Read(to []byte) (int, error) {
|
||||
return n - 4, nil
|
||||
}
|
||||
|
||||
// Write pushes one IP packet onto the utun device.
|
||||
// Write pushes one IP packet onto the utun device. Safe for concurrent use:
|
||||
// the AF prefix and iovecs are per-call stack state, and the fd write itself
|
||||
// serializes on the runtime's fd mutex (see the Queue contract in tio.go).
|
||||
func (t *tun) Write(from []byte) (int, error) {
|
||||
if len(from) == 0 {
|
||||
return 0, syscall.EIO
|
||||
@@ -606,10 +608,6 @@ func (t *tun) Name() string {
|
||||
return t.Device
|
||||
}
|
||||
|
||||
func (t *tun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for darwin")
|
||||
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||
}
|
||||
|
||||
+26
-24
@@ -10,6 +10,7 @@ import (
|
||||
|
||||
"github.com/rcrowley/go-metrics"
|
||||
"github.com/slackhq/nebula/iputil"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
)
|
||||
|
||||
@@ -23,6 +24,23 @@ type disabledTun struct {
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
// Read hands the next queued packet to a reader, copying it into b. Reads
|
||||
// from concurrent queues are safe: the channel receive serializes them and
|
||||
// each queue copies into its own private scratch buffer.
|
||||
func (t *disabledTun) Read(b []byte) (int, error) {
|
||||
r, ok := <-t.read
|
||||
if !ok {
|
||||
return 0, io.EOF
|
||||
}
|
||||
|
||||
t.tx.Inc(1)
|
||||
if t.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
t.l.Debug("Write payload", "raw", prettyPacket(r))
|
||||
}
|
||||
|
||||
return copy(b, r), nil
|
||||
}
|
||||
|
||||
func newDisabledTun(vpnNetworks []netip.Prefix, queueLen int, metricsEnabled bool, l *slog.Logger) *disabledTun {
|
||||
tun := &disabledTun{
|
||||
vpnNetworks: vpnNetworks,
|
||||
@@ -57,24 +75,6 @@ func (*disabledTun) Name() string {
|
||||
return "disabled"
|
||||
}
|
||||
|
||||
func (t *disabledTun) Read(b []byte) (int, error) {
|
||||
r, ok := <-t.read
|
||||
if !ok {
|
||||
return 0, io.EOF
|
||||
}
|
||||
|
||||
if len(r) > len(b) {
|
||||
return 0, fmt.Errorf("packet larger than mtu: %d > %d bytes", len(r), len(b))
|
||||
}
|
||||
|
||||
t.tx.Inc(1)
|
||||
if t.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
t.l.Debug("Write payload", "raw", prettyPacket(r))
|
||||
}
|
||||
|
||||
return copy(b, r), nil
|
||||
}
|
||||
|
||||
func (t *disabledTun) handleICMPEchoRequest(b []byte) bool {
|
||||
out := make([]byte, len(b))
|
||||
out = iputil.CreateICMPEchoResponse(b, out)
|
||||
@@ -106,12 +106,14 @@ func (t *disabledTun) Write(b []byte) (int, error) {
|
||||
return len(b), nil
|
||||
}
|
||||
|
||||
func (t *disabledTun) SupportsMultiqueue() bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func (t *disabledTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return t, nil
|
||||
func (t *disabledTun) Queues(n int) ([]tio.Queue, error) {
|
||||
out := make([]tio.Queue, n)
|
||||
for i := range out {
|
||||
// NoClose: the shared channel and metrics are owned by the
|
||||
// disabledTun; Close on the device tears them down once for everybody.
|
||||
out[i] = tio.NewSingleQueueNoClose(t, defaultBatchBufSize)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (t *disabledTun) Close() error {
|
||||
|
||||
@@ -1,120 +0,0 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
// +build linux,!android,!e2e_testing
|
||||
|
||||
package overlay
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// newReadPipe returns a read fd. The matching write fd is registered for cleanup.
|
||||
// The caller takes ownership of the read fd (pass it to newTunFd / newFriend).
|
||||
func newReadPipe(t *testing.T) int {
|
||||
t.Helper()
|
||||
var fds [2]int
|
||||
if err := unix.Pipe2(fds[:], unix.O_CLOEXEC); err != nil {
|
||||
t.Fatalf("pipe2: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = unix.Close(fds[1]) })
|
||||
return fds[0]
|
||||
}
|
||||
|
||||
func TestTunFile_WakeForShutdown_UnblocksRead(t *testing.T) {
|
||||
tf, err := newTunFd(newReadPipe(t))
|
||||
if err != nil {
|
||||
t.Fatalf("newTunFd: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = tf.Close() })
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := tf.Read(make([]byte, 64))
|
||||
done <- err
|
||||
}()
|
||||
|
||||
// Verify Read is actually blocked in poll.
|
||||
select {
|
||||
case err := <-done:
|
||||
t.Fatalf("Read returned before shutdown signal: %v", err)
|
||||
case <-time.After(50 * time.Millisecond):
|
||||
}
|
||||
|
||||
if err := tf.wakeForShutdown(); err != nil {
|
||||
t.Fatalf("wakeForShutdown: %v", err)
|
||||
}
|
||||
|
||||
select {
|
||||
case err := <-done:
|
||||
if !errors.Is(err, os.ErrClosed) {
|
||||
t.Fatalf("expected os.ErrClosed, got %v", err)
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("Read did not wake on shutdown")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunFile_WakeForShutdown_WakesFriends(t *testing.T) {
|
||||
parent, err := newTunFd(newReadPipe(t))
|
||||
if err != nil {
|
||||
t.Fatalf("newTunFd: %v", err)
|
||||
}
|
||||
friend, err := parent.newFriend(newReadPipe(t))
|
||||
if err != nil {
|
||||
_ = parent.Close()
|
||||
t.Fatalf("newFriend: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = friend.Close()
|
||||
_ = parent.Close()
|
||||
})
|
||||
|
||||
readers := []*tunFile{parent, friend}
|
||||
errs := make([]error, len(readers))
|
||||
var wg sync.WaitGroup
|
||||
for i, r := range readers {
|
||||
wg.Add(1)
|
||||
go func(i int, r *tunFile) {
|
||||
defer wg.Done()
|
||||
_, errs[i] = r.Read(make([]byte, 64))
|
||||
}(i, r)
|
||||
}
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
if err := parent.wakeForShutdown(); err != nil {
|
||||
t.Fatalf("wakeForShutdown: %v", err)
|
||||
}
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() { wg.Wait(); close(done) }()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("readers did not wake")
|
||||
}
|
||||
|
||||
for i, err := range errs {
|
||||
if !errors.Is(err, os.ErrClosed) {
|
||||
t.Errorf("reader %d: expected os.ErrClosed, got %v", i, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunFile_Close_Idempotent(t *testing.T) {
|
||||
tf, err := newTunFd(newReadPipe(t))
|
||||
if err != nil {
|
||||
t.Fatalf("newTunFd: %v", err)
|
||||
}
|
||||
if err := tf.Close(); err != nil {
|
||||
t.Fatalf("first Close: %v", err)
|
||||
}
|
||||
if err := tf.Close(); err != nil {
|
||||
t.Fatalf("second Close should be a no-op, got %v", err)
|
||||
}
|
||||
}
|
||||
@@ -7,7 +7,6 @@ import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
@@ -20,7 +19,7 @@ import (
|
||||
"github.com/gaissmai/bart"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
netroute "golang.org/x/net/route"
|
||||
@@ -561,12 +560,8 @@ func (t *tun) Name() string {
|
||||
return t.Device
|
||||
}
|
||||
|
||||
func (t *tun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for freebsd")
|
||||
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||
}
|
||||
|
||||
func (t *tun) addRoutes(logErrors bool) error {
|
||||
|
||||
+3
-6
@@ -16,6 +16,7 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
"golang.org/x/sys/unix"
|
||||
@@ -159,10 +160,6 @@ func (t *tun) Name() string {
|
||||
return "iOS"
|
||||
}
|
||||
|
||||
func (t *tun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for ios")
|
||||
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||
}
|
||||
|
||||
+166
-257
@@ -4,10 +4,8 @@
|
||||
package overlay
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
@@ -20,188 +18,25 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
"github.com/vishvananda/netlink"
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// tunFile wraps a TUN file descriptor with poll-based reads. The FD provided will be changed to non-blocking.
|
||||
// A shared eventfd allows Close to wake all readers blocked in poll.
|
||||
type tunFile struct {
|
||||
fd int
|
||||
shutdownFd int
|
||||
lastOne bool
|
||||
readPoll [2]unix.PollFd
|
||||
writePoll [2]unix.PollFd
|
||||
closed bool
|
||||
}
|
||||
|
||||
// newFriend makes a tunFile for a MultiQueueReader that copies the shutdown eventfd from the parent tun
|
||||
func (r *tunFile) newFriend(fd int) (*tunFile, error) {
|
||||
if err := unix.SetNonblock(fd, true); err != nil {
|
||||
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
||||
}
|
||||
return &tunFile{
|
||||
fd: fd,
|
||||
shutdownFd: r.shutdownFd,
|
||||
readPoll: [2]unix.PollFd{
|
||||
{Fd: int32(fd), Events: unix.POLLIN},
|
||||
{Fd: int32(r.shutdownFd), Events: unix.POLLIN},
|
||||
},
|
||||
writePoll: [2]unix.PollFd{
|
||||
{Fd: int32(fd), Events: unix.POLLOUT},
|
||||
{Fd: int32(r.shutdownFd), Events: unix.POLLIN},
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func newTunFd(fd int) (*tunFile, error) {
|
||||
if err := unix.SetNonblock(fd, true); err != nil {
|
||||
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
||||
}
|
||||
|
||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||
}
|
||||
|
||||
out := &tunFile{
|
||||
fd: fd,
|
||||
shutdownFd: shutdownFd,
|
||||
lastOne: true,
|
||||
readPoll: [2]unix.PollFd{
|
||||
{Fd: int32(fd), Events: unix.POLLIN},
|
||||
{Fd: int32(shutdownFd), Events: unix.POLLIN},
|
||||
},
|
||||
writePoll: [2]unix.PollFd{
|
||||
{Fd: int32(fd), Events: unix.POLLOUT},
|
||||
{Fd: int32(shutdownFd), Events: unix.POLLIN},
|
||||
},
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (r *tunFile) blockOnRead() error {
|
||||
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
||||
var err error
|
||||
for {
|
||||
_, err = unix.Poll(r.readPoll[:], -1)
|
||||
if err != unix.EINTR {
|
||||
break
|
||||
}
|
||||
}
|
||||
//always reset these!
|
||||
tunEvents := r.readPoll[0].Revents
|
||||
shutdownEvents := r.readPoll[1].Revents
|
||||
r.readPoll[0].Revents = 0
|
||||
r.readPoll[1].Revents = 0
|
||||
//do the err check before trusting the potentially bogus bits we just got
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
||||
return os.ErrClosed
|
||||
} else if tunEvents&problemFlags != 0 {
|
||||
return os.ErrClosed
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *tunFile) blockOnWrite() error {
|
||||
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
||||
var err error
|
||||
for {
|
||||
_, err = unix.Poll(r.writePoll[:], -1)
|
||||
if err != unix.EINTR {
|
||||
break
|
||||
}
|
||||
}
|
||||
//always reset these!
|
||||
tunEvents := r.writePoll[0].Revents
|
||||
shutdownEvents := r.writePoll[1].Revents
|
||||
r.writePoll[0].Revents = 0
|
||||
r.writePoll[1].Revents = 0
|
||||
//do the err check before trusting the potentially bogus bits we just got
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
||||
return os.ErrClosed
|
||||
} else if tunEvents&problemFlags != 0 {
|
||||
return os.ErrClosed
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *tunFile) Read(buf []byte) (int, error) {
|
||||
for {
|
||||
if n, err := unix.Read(r.fd, buf); err == nil {
|
||||
return n, nil
|
||||
} else if err == unix.EAGAIN {
|
||||
if err = r.blockOnRead(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
continue
|
||||
} else if err == unix.EINTR {
|
||||
continue
|
||||
} else if err == unix.EBADF {
|
||||
return 0, os.ErrClosed
|
||||
} else {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (r *tunFile) Write(buf []byte) (int, error) {
|
||||
for {
|
||||
if n, err := unix.Write(r.fd, buf); err == nil {
|
||||
return n, nil
|
||||
} else if err == unix.EAGAIN {
|
||||
if err = r.blockOnWrite(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
continue
|
||||
} else if err == unix.EINTR {
|
||||
continue
|
||||
} else if err == unix.EBADF {
|
||||
return 0, os.ErrClosed
|
||||
} else {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (r *tunFile) wakeForShutdown() error {
|
||||
var buf [8]byte
|
||||
binary.NativeEndian.PutUint64(buf[:], 1)
|
||||
_, err := unix.Write(int(r.readPoll[1].Fd), buf[:])
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *tunFile) Close() error {
|
||||
if r.closed { // avoid closing more than once. Technically a fd could get re-used, which would be a problem
|
||||
return nil
|
||||
}
|
||||
r.closed = true
|
||||
if r.lastOne {
|
||||
_ = unix.Close(r.shutdownFd)
|
||||
}
|
||||
return unix.Close(r.fd)
|
||||
}
|
||||
|
||||
type tun struct {
|
||||
*tunFile
|
||||
readers []*tunFile
|
||||
closeLock sync.Mutex
|
||||
Device string
|
||||
vpnNetworks []netip.Prefix
|
||||
MaxMTU int
|
||||
DefaultMTU int
|
||||
TXQueueLen int
|
||||
deviceIndex int
|
||||
ioctlFd uintptr
|
||||
readers tio.QueueSet
|
||||
closeLock sync.Mutex
|
||||
Device string
|
||||
vpnNetworks []netip.Prefix
|
||||
MaxMTU int
|
||||
DefaultMTU int
|
||||
TXQueueLen int
|
||||
deviceIndex int
|
||||
ioctlFd uintptr
|
||||
vnetHdr bool
|
||||
offloadFlags uint
|
||||
|
||||
Routes atomic.Pointer[[]Route]
|
||||
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
||||
@@ -240,56 +75,112 @@ type ifreqQLEN struct {
|
||||
}
|
||||
|
||||
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||
t, err := newTunGeneric(c, l, deviceFd, vpnNetworks)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
// We don't know what flags the caller opened this fd with and can't turn
|
||||
// on IFF_VNET_HDR after TUNSETIFF, so skip offload on inherited fds.
|
||||
return newTunGeneric(c, l, deviceFd, false, 0, vpnNetworks, "tun0")
|
||||
}
|
||||
|
||||
// openTunDev opens /dev/net/tun, creating the device node first if it's
|
||||
// missing (docker containers occasionally omit it).
|
||||
func openTunDev() (int, error) {
|
||||
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||
if err == nil {
|
||||
return fd, nil
|
||||
}
|
||||
if !os.IsNotExist(err) {
|
||||
return -1, err
|
||||
}
|
||||
if err = os.MkdirAll("/dev/net", 0755); err != nil {
|
||||
return -1, fmt.Errorf("/dev/net/tun doesn't exist, failed to mkdir -p /dev/net: %w", err)
|
||||
}
|
||||
if err = unix.Mknod("/dev/net/tun", unix.S_IFCHR|0600, int(unix.Mkdev(10, 200))); err != nil {
|
||||
return -1, fmt.Errorf("failed to create /dev/net/tun: %w", err)
|
||||
}
|
||||
fd, err = unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||
if err != nil {
|
||||
return -1, fmt.Errorf("created /dev/net/tun, but still failed: %w", err)
|
||||
}
|
||||
return fd, nil
|
||||
}
|
||||
|
||||
t.Device = "tun0"
|
||||
// tunSetIff runs TUNSETIFF with the given flags and returns the kernel-chosen device name on success.
|
||||
func tunSetIff(fd int, name string, flags uint16) (string, error) {
|
||||
var req ifReq
|
||||
req.Flags = flags
|
||||
copy(req.Name[:], name)
|
||||
if err := ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return strings.Trim(string(req.Name[:]), "\x00"), nil
|
||||
}
|
||||
|
||||
return t, nil
|
||||
// tsoOffloadFlags are the TUN_F_* bits we ask the kernel to enable when a TSO-capable TUN is available.
|
||||
const tsoOffloadFlags = unix.TUN_F_CSUM | unix.TUN_F_TSO4 | unix.TUN_F_TSO6 | unix.TUN_F_TSO_ECN
|
||||
|
||||
// usoAndTSOOffloadFlags adds UDP Segmentation Offload to tsoOffloadFlags.
|
||||
// Requires Linux >= 6.2; older kernels reject it and we fall back to TCP-only TSO
|
||||
const usoAndTSOOffloadFlags = tsoOffloadFlags | unix.TUN_F_USO4 | unix.TUN_F_USO6
|
||||
|
||||
func offloadUSOEnabled(offloadFlags uint) bool {
|
||||
return offloadFlags&(unix.TUN_F_USO4|unix.TUN_F_USO6) != 0
|
||||
}
|
||||
|
||||
func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) {
|
||||
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||
if err != nil {
|
||||
// If /dev/net/tun doesn't exist, try to create it (will happen in docker)
|
||||
if os.IsNotExist(err) {
|
||||
err = os.MkdirAll("/dev/net", 0755)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("/dev/net/tun doesn't exist, failed to mkdir -p /dev/net: %w", err)
|
||||
}
|
||||
err = unix.Mknod("/dev/net/tun", unix.S_IFCHR|0600, int(unix.Mkdev(10, 200)))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create /dev/net/tun: %w", err)
|
||||
}
|
||||
|
||||
fd, err = unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("created /dev/net/tun, but still failed: %w", err)
|
||||
}
|
||||
} else {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
var req ifReq
|
||||
req.Flags = uint16(unix.IFF_TUN | unix.IFF_NO_PI)
|
||||
var err error
|
||||
// IFF_TUN_EXCL prevents us from attaching to an already-running tun
|
||||
baseFlags := uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_TUN_EXCL)
|
||||
if multiqueue {
|
||||
req.Flags |= unix.IFF_MULTI_QUEUE
|
||||
baseFlags |= unix.IFF_MULTI_QUEUE
|
||||
}
|
||||
nameStr := c.GetString("tun.dev", "")
|
||||
copy(req.Name[:], nameStr)
|
||||
if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return nil, &NameError{
|
||||
Name: nameStr,
|
||||
Underlying: err,
|
||||
useOffloads := c.GetBool("tun.use_offloads", true)
|
||||
|
||||
var fd int
|
||||
var name string
|
||||
var offloadFlags uint
|
||||
if useOffloads {
|
||||
fd, err = openTunDev()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// First try to enable IFF_VNET_HDR via TUNSETIFF and negotiate TUN_F_* offloads
|
||||
// We try TSO+USO first, fall back to TSO-only on kernels without USO (Linux < 6.2),
|
||||
// and finally give up on virtio headers entirely and reopen as a plain TUN if neither offload mask is accepted.
|
||||
|
||||
// offloadFlags is the exact TUN_F_* mask the kernel accepted.
|
||||
// We save it so addQueue can replay the identical device-wide mask on added queues
|
||||
name, err = tunSetIff(fd, nameStr, baseFlags|unix.IFF_VNET_HDR)
|
||||
if err != nil {
|
||||
_ = unix.Close(fd)
|
||||
useOffloads = false
|
||||
} else {
|
||||
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(usoAndTSOOffloadFlags)); err == nil {
|
||||
offloadFlags = usoAndTSOOffloadFlags
|
||||
} else if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(tsoOffloadFlags)); err == nil {
|
||||
offloadFlags = tsoOffloadFlags
|
||||
} else {
|
||||
l.Warn("Failed to enable TUN offload (TSO); proceeding without virtio headers", "error", err)
|
||||
_ = unix.Close(fd)
|
||||
useOffloads = false
|
||||
}
|
||||
}
|
||||
}
|
||||
name := strings.Trim(string(req.Name[:]), "\x00")
|
||||
|
||||
t, err := newTunGeneric(c, l, fd, vpnNetworks)
|
||||
if !useOffloads {
|
||||
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}
|
||||
}
|
||||
}
|
||||
|
||||
l.Info("TUN offload status", "tso", useOffloads, "uso", offloadUSOEnabled(offloadFlags))
|
||||
|
||||
t, err := newTunGeneric(c, l, fd, useOffloads, offloadFlags, vpnNetworks, name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -299,17 +190,37 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
|
||||
return t, nil
|
||||
}
|
||||
|
||||
// newTunGeneric does all the stuff common to different tun initialization paths. It will close your files on error.
|
||||
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||
tfd, err := newTunFd(fd)
|
||||
// newTunGeneric does all the stuff common to different tun initialization paths.
|
||||
// It will close your files on error.
|
||||
// offloadFlags is the TUN_F_* mask newTun negotiated (ignored when vnetHdr is false)
|
||||
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFlags uint, vpnNetworks []netip.Prefix, name string) (*tun, error) {
|
||||
var qs tio.QueueSet
|
||||
var err error
|
||||
if vnetHdr {
|
||||
qs, err = tio.NewOffloadQueueSet(offloadUSOEnabled(offloadFlags), l)
|
||||
} else {
|
||||
qs, err = tio.NewPollQueueSet()
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return nil, err
|
||||
}
|
||||
err = qs.Add(fd)
|
||||
if err != nil {
|
||||
// Add only appends on success, so closing the set here can't
|
||||
// double-close fd; it releases the set's shutdown eventfd.
|
||||
_ = unix.Close(fd)
|
||||
_ = qs.Close()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
t := &tun{
|
||||
tunFile: tfd,
|
||||
readers: []*tunFile{tfd},
|
||||
Device: name,
|
||||
readers: qs,
|
||||
closeLock: sync.Mutex{},
|
||||
vnetHdr: vnetHdr,
|
||||
offloadFlags: offloadFlags,
|
||||
vpnNetworks: vpnNetworks,
|
||||
TXQueueLen: c.GetInt("tun.tx_queue", 500),
|
||||
useSystemRoutes: c.GetBool("tun.use_system_route_table", false),
|
||||
@@ -407,36 +318,49 @@ func (t *tun) reload(c *config.C, initial bool) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *tun) SupportsMultiqueue() bool {
|
||||
return true
|
||||
// Queues opens additional kernel multiqueue fds until the device has n queues, then returns them all.
|
||||
func (t *tun) Queues(n int) ([]tio.Queue, error) {
|
||||
for len(t.readers.Queues()) < n {
|
||||
if err := t.addQueue(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return t.readers.Queues(), nil
|
||||
}
|
||||
|
||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
// addQueue opens one more IFF_MULTI_QUEUE fd on the device and adds it to the queue set.
|
||||
func (t *tun) addQueue() error {
|
||||
t.closeLock.Lock()
|
||||
defer t.closeLock.Unlock()
|
||||
|
||||
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return err
|
||||
}
|
||||
|
||||
var req ifReq
|
||||
req.Flags = uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_MULTI_QUEUE)
|
||||
copy(req.Name[:], t.Device)
|
||||
if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
||||
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 {
|
||||
_ = unix.Close(fd)
|
||||
return nil, err
|
||||
return err
|
||||
}
|
||||
|
||||
out, err := t.tunFile.newFriend(fd)
|
||||
if t.vnetHdr {
|
||||
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(t.offloadFlags)); err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return fmt.Errorf("failed to enable offload on multiqueue tun fd: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
err = t.readers.Add(fd)
|
||||
if err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return nil, err
|
||||
return err
|
||||
}
|
||||
|
||||
t.readers = append(t.readers, out)
|
||||
|
||||
return out, nil
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
||||
@@ -613,6 +537,13 @@ func (t *tun) setDefaultRoute(cidr netip.Prefix) error {
|
||||
Table: unix.RT_TABLE_MAIN,
|
||||
Type: unix.RTN_UNICAST,
|
||||
}
|
||||
// Match the metric the kernel uses for its auto-installed connected route,
|
||||
// so RouteReplace overwrites it in place instead of adding a second route at a worse metric.
|
||||
// IPv6 connected routes are installed at metric 256 (IP6_RT_PRIO_KERN); IPv4 uses 0.
|
||||
// Without this, the kernel route wins lookups and our MTU / AdvMSS / Features never apply on v6.
|
||||
if cidr.Addr().Is6() {
|
||||
nr.Priority = 256
|
||||
}
|
||||
err := netlink.RouteReplace(&nr)
|
||||
if err != nil {
|
||||
t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr)
|
||||
@@ -888,32 +819,10 @@ func (t *tun) Close() error {
|
||||
t.routeChan = nil
|
||||
}
|
||||
|
||||
// Signal all readers blocked in poll to wake up and exit
|
||||
_ = t.tunFile.wakeForShutdown()
|
||||
|
||||
if t.ioctlFd > 0 {
|
||||
_ = unix.Close(int(t.ioctlFd))
|
||||
t.ioctlFd = 0
|
||||
}
|
||||
|
||||
for i := range t.readers {
|
||||
if i == 0 {
|
||||
continue //we want to close the zeroth reader last
|
||||
}
|
||||
err := t.readers[i].Close()
|
||||
if err != nil {
|
||||
t.l.Error("error closing tun reader", "reader", i, "error", err)
|
||||
} else {
|
||||
t.l.Info("closed tun reader", "reader", i)
|
||||
}
|
||||
}
|
||||
|
||||
//this is t.readers[0] too
|
||||
err := t.tunFile.Close()
|
||||
if err != nil {
|
||||
t.l.Error("error closing tun reader", "reader", 0, "error", err)
|
||||
} else {
|
||||
t.l.Info("closed tun reader", "reader", 0)
|
||||
}
|
||||
return err
|
||||
return t.readers.Close()
|
||||
}
|
||||
|
||||
@@ -3,7 +3,9 @@
|
||||
|
||||
package overlay
|
||||
|
||||
import "testing"
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
var runAdvMSSTests = []struct {
|
||||
name string
|
||||
@@ -32,3 +34,65 @@ func TestTunAdvMSS(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestOffloadUSOEnabled pins the single source of truth for the per-queue USO
|
||||
// capability: it is derived from the negotiated offload mask, so the mask
|
||||
// stored on the tun and the capability reported to coalescers cannot drift.
|
||||
func TestOffloadUSOEnabled(t *testing.T) {
|
||||
// usoAndTSOOffloadFlags must be a strict superset of tsoOffloadFlags. Otherwise
|
||||
// the TSO-only fallback (and the historic hardcoded-mask bug in
|
||||
// addQueue) would not actually be a downgrade.
|
||||
if usoAndTSOOffloadFlags&tsoOffloadFlags != tsoOffloadFlags {
|
||||
t.Fatalf("usoAndTSOOffloadFlags (%#x) is not a superset of tsoOffloadFlags (%#x)", usoAndTSOOffloadFlags, tsoOffloadFlags)
|
||||
}
|
||||
if usoAndTSOOffloadFlags == tsoOffloadFlags {
|
||||
t.Fatal("usoAndTSOOffloadFlags must add bits beyond tsoOffloadFlags")
|
||||
}
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
offloadFlags uint
|
||||
wantUSO bool
|
||||
}{
|
||||
{"uso-negotiated", usoAndTSOOffloadFlags, true},
|
||||
{"tso-fallback", tsoOffloadFlags, false},
|
||||
{"no-vnet-hdr", 0, false},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := offloadUSOEnabled(tc.offloadFlags); got != tc.wantUSO {
|
||||
t.Fatalf("offloadUSOEnabled(%#x) = %v, want %v", tc.offloadFlags, got, tc.wantUSO)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestAddQueueReplaysNegotiatedMask guards the device-wide TUNSETOFFLOAD
|
||||
// downgrade bug: addQueue must issue the exact mask newTun negotiated
|
||||
// (t.offloadFlags), not a hardcoded TSO-only mask. Because TUNSETOFFLOAD is
|
||||
// per-netdev, a narrower mask on an added queue silently disables USO for
|
||||
// every queue on a USO-capable kernel while the queues keep advertising it.
|
||||
//
|
||||
// A full multi-queue exercise needs /dev/net/tun and CAP_NET_ADMIN, which are
|
||||
// not available in CI/sandbox, so this asserts on the struct field that the
|
||||
// TUNSETOFFLOAD argument is read from.
|
||||
func TestAddQueueReplaysNegotiatedMask(t *testing.T) {
|
||||
t.Run("uso-negotiated", func(t *testing.T) {
|
||||
tn := &tun{vnetHdr: true, offloadFlags: usoAndTSOOffloadFlags}
|
||||
// The ioctl argument in addQueue is uintptr(t.offloadFlags);
|
||||
// it must equal the negotiated USO mask, and must NOT be the TSO-only
|
||||
// mask (the original bug).
|
||||
if tn.offloadFlags != usoAndTSOOffloadFlags {
|
||||
t.Fatalf("offloadFlags = %#x, want %#x", tn.offloadFlags, usoAndTSOOffloadFlags)
|
||||
}
|
||||
if tn.offloadFlags == tsoOffloadFlags {
|
||||
t.Fatal("added queue would downgrade USO: offloadFlags must not be the TSO-only mask when USO was negotiated")
|
||||
}
|
||||
})
|
||||
t.Run("tso-fallback", func(t *testing.T) {
|
||||
tn := &tun{vnetHdr: true, offloadFlags: tsoOffloadFlags}
|
||||
if tn.offloadFlags != tsoOffloadFlags {
|
||||
t.Fatalf("offloadFlags = %#x, want %#x", tn.offloadFlags, tsoOffloadFlags)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -6,7 +6,6 @@ package overlay
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"os"
|
||||
@@ -17,6 +16,7 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
netroute "golang.org/x/net/route"
|
||||
@@ -390,12 +390,8 @@ func (t *tun) Name() string {
|
||||
return t.Device
|
||||
}
|
||||
|
||||
func (t *tun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for netbsd")
|
||||
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||
}
|
||||
|
||||
func (t *tun) addRoutes(logErrors bool) error {
|
||||
|
||||
@@ -6,7 +6,6 @@ package overlay
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"os"
|
||||
@@ -17,6 +16,7 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
netroute "golang.org/x/net/route"
|
||||
@@ -138,8 +138,8 @@ func tunWritev(fd int, iovecs []unix.Iovec) (n int, err error)
|
||||
//go:noescape
|
||||
func tunReadv(fd int, iovecs []unix.Iovec) (n int, err error)
|
||||
|
||||
// Read pulls one IP packet off the tun device, scattering the 4 byte protocol header away from the
|
||||
// packet so the payload lands directly in to.
|
||||
// Read pulls one IP packet off the tun device, scattering the 4 byte protocol header away from
|
||||
// the packet so the payload lands directly in to.
|
||||
func (t *tun) Read(to []byte) (int, error) {
|
||||
var head [4]byte
|
||||
|
||||
@@ -369,12 +369,8 @@ func (t *tun) Name() string {
|
||||
return t.Device
|
||||
}
|
||||
|
||||
func (t *tun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for openbsd")
|
||||
func (t *tun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||
}
|
||||
|
||||
func (t *tun) addRoutes(logErrors bool) error {
|
||||
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
)
|
||||
@@ -177,10 +178,6 @@ func (t *TestTun) Read(b []byte) (int, error) {
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func (t *TestTun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *TestTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented")
|
||||
func (t *TestTun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, udp.MTU)}, nil
|
||||
}
|
||||
|
||||
+7
-11
@@ -6,7 +6,6 @@ package overlay
|
||||
import (
|
||||
"crypto"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"os"
|
||||
@@ -18,6 +17,7 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
"github.com/slackhq/nebula/wintun"
|
||||
@@ -47,6 +47,10 @@ type winTun struct {
|
||||
tun *wintun.NativeTun
|
||||
}
|
||||
|
||||
func (t *winTun) Read(b []byte) (int, error) {
|
||||
return t.tun.Read(b, 0)
|
||||
}
|
||||
|
||||
func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (Device, error) {
|
||||
return nil, fmt.Errorf("newTunFromFd not supported in Windows")
|
||||
}
|
||||
@@ -255,20 +259,12 @@ func (t *winTun) Name() string {
|
||||
return t.Device
|
||||
}
|
||||
|
||||
func (t *winTun) Read(b []byte) (int, error) {
|
||||
return t.tun.Read(b, 0)
|
||||
}
|
||||
|
||||
func (t *winTun) Write(b []byte) (int, error) {
|
||||
return t.tun.Write(b, 0)
|
||||
}
|
||||
|
||||
func (t *winTun) SupportsMultiqueue() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (t *winTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for windows")
|
||||
func (t *winTun) Queues(int) ([]tio.Queue, error) {
|
||||
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil
|
||||
}
|
||||
|
||||
func (t *winTun) Close() error {
|
||||
|
||||
+13
-6
@@ -6,6 +6,7 @@ import (
|
||||
"net/netip"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
)
|
||||
|
||||
@@ -46,12 +47,16 @@ func (d *UserDevice) RoutesFor(ip netip.Addr) routing.Gateways {
|
||||
return routing.Gateways{routing.NewGateway(ip, 1)}
|
||||
}
|
||||
|
||||
func (d *UserDevice) SupportsMultiqueue() bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func (d *UserDevice) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return d, nil
|
||||
func (d *UserDevice) Queues(n int) ([]tio.Queue, error) {
|
||||
out := make([]tio.Queue, n)
|
||||
for i := range out {
|
||||
// All queues share the underlying pipes (the io.Pipe serializes
|
||||
// concurrent callers) but each owns a private scratch buffer so
|
||||
// concurrent Reads across queues never alias. NoClose: the pipes are
|
||||
// owned by the UserDevice and torn down once by UserDevice.Close.
|
||||
out[i] = tio.NewSingleQueueNoClose(d, defaultBatchBufSize)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (d *UserDevice) Pipe() (*io.PipeReader, *io.PipeWriter) {
|
||||
@@ -61,9 +66,11 @@ func (d *UserDevice) Pipe() (*io.PipeReader, *io.PipeWriter) {
|
||||
func (d *UserDevice) Read(p []byte) (n int, err error) {
|
||||
return d.outboundReader.Read(p)
|
||||
}
|
||||
|
||||
func (d *UserDevice) Write(p []byte) (n int, err error) {
|
||||
return d.inboundWriter.Write(p)
|
||||
}
|
||||
|
||||
func (d *UserDevice) Close() error {
|
||||
d.inboundWriter.Close()
|
||||
d.outboundWriter.Close()
|
||||
|
||||
@@ -0,0 +1,163 @@
|
||||
package overlay
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
)
|
||||
|
||||
// newTestUserDevice returns the concrete *UserDevice so tests can reach Pipe()
|
||||
// and the internal queue plumbing.
|
||||
func newTestUserDevice(t *testing.T) *UserDevice {
|
||||
t.Helper()
|
||||
dev, err := NewUserDevice([]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||
if err != nil {
|
||||
t.Fatalf("NewUserDevice: %v", err)
|
||||
}
|
||||
ud, ok := dev.(*UserDevice)
|
||||
if !ok {
|
||||
t.Fatalf("NewUserDevice returned %T, want *UserDevice", dev)
|
||||
}
|
||||
return ud
|
||||
}
|
||||
|
||||
// TestUserDeviceReadersDistinctBuffers ensures each Queue is actually different
|
||||
func TestUserDeviceReadersDistinctBuffers(t *testing.T) {
|
||||
d := newTestUserDevice(t)
|
||||
|
||||
readers, err := d.Queues(2)
|
||||
if err != nil {
|
||||
t.Fatalf("Queues: %v", err)
|
||||
}
|
||||
if len(readers) != 2 {
|
||||
t.Fatalf("Queues(2) returned %d queues, want 2", len(readers))
|
||||
}
|
||||
|
||||
// Distinct queue objects.
|
||||
if readers[0] == readers[1] {
|
||||
t.Fatal("Queues(2) returned the same queue object twice")
|
||||
}
|
||||
|
||||
// Drive one packet through each queue and confirm the borrowed bytes from
|
||||
// the first read are NOT clobbered by the second read. With a shared
|
||||
// buffer, reading pkt1 into q1 would corrupt q0's still-borrowed slice.
|
||||
_, ow := d.Pipe()
|
||||
|
||||
pkt0 := []byte("packet-zero-aaaaaaaa")
|
||||
pkt1 := []byte("packet-one-bbbbbbbbb")
|
||||
|
||||
// The pipe is unbuffered, so writes block until a reader consumes them.
|
||||
// Serialize: write pkt0 (read on q0), then write pkt1 (read on q1).
|
||||
go func() {
|
||||
if _, err := ow.Write(pkt0); err != nil {
|
||||
t.Errorf("write pkt0: %v", err)
|
||||
}
|
||||
if _, err := ow.Write(pkt1); err != nil {
|
||||
t.Errorf("write pkt1: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
got0, err := readers[0].Read()
|
||||
if err != nil {
|
||||
t.Fatalf("q0.Read: %v", err)
|
||||
}
|
||||
if len(got0) != 1 || string(got0[0].Bytes) != string(pkt0) {
|
||||
t.Fatalf("q0 first read = %q, want %q", firstBytes(got0), pkt0)
|
||||
}
|
||||
// Hold onto q0's borrowed slice across q1's read.
|
||||
borrowed := got0[0].Bytes
|
||||
|
||||
got1, err := readers[1].Read()
|
||||
if err != nil {
|
||||
t.Fatalf("q1.Read: %v", err)
|
||||
}
|
||||
if len(got1) != 1 || string(got1[0].Bytes) != string(pkt1) {
|
||||
t.Fatalf("q1 read = %q, want %q", firstBytes(got1), pkt1)
|
||||
}
|
||||
|
||||
// q0's borrowed bytes must still hold pkt0 - a shared buffer would now
|
||||
// show pkt1's contents.
|
||||
if string(borrowed) != string(pkt0) {
|
||||
t.Fatalf("q0 borrowed bytes were clobbered by q1's read: got %q, want %q", borrowed, pkt0)
|
||||
}
|
||||
}
|
||||
|
||||
// TestUserDeviceReadersConcurrentRace exercises two queues reading distinct
|
||||
// packets concurrently. Run it under `go test -race`: with the old
|
||||
// shared-buffer implementation the concurrent Reads raced on readBuf/batchRet
|
||||
// and corrupted each other's returned slices.
|
||||
func TestUserDeviceReadersConcurrentRace(t *testing.T) {
|
||||
d := newTestUserDevice(t)
|
||||
readers, err := d.Queues(2)
|
||||
if err != nil {
|
||||
t.Fatalf("Queues: %v", err)
|
||||
}
|
||||
_, ow := d.Pipe()
|
||||
|
||||
const iterations = 200
|
||||
|
||||
errs := make(chan error, 3)
|
||||
|
||||
// Each reader parks in Read on the shared outboundReader; io.Pipe hands
|
||||
// each write to whichever reader is currently waiting. We only care that
|
||||
// concurrent Reads into distinct buffers are race-free, so any parked
|
||||
// reader may serve any write.
|
||||
var wg sync.WaitGroup
|
||||
run := func(idx int) {
|
||||
defer wg.Done()
|
||||
for i := 0; i < iterations; i++ {
|
||||
pkts, err := readers[idx].Read()
|
||||
if err != nil {
|
||||
errs <- err
|
||||
return
|
||||
}
|
||||
if len(pkts) != 1 {
|
||||
errs <- fmt.Errorf("reader %d: got %d packets, want 1", idx, len(pkts))
|
||||
return
|
||||
}
|
||||
// Touch every byte of the borrowed slice while the other reader
|
||||
// may be mid-Read; a shared buffer would race here.
|
||||
total := 0
|
||||
for _, c := range pkts[0].Bytes {
|
||||
total += int(c)
|
||||
}
|
||||
_ = total
|
||||
}
|
||||
}
|
||||
|
||||
wg.Add(2)
|
||||
go run(0)
|
||||
go run(1)
|
||||
|
||||
// Feed 2*iterations packets. io.Pipe copies each write straight into the
|
||||
// waiting reader's private buffer, so reusing buf between writes is safe.
|
||||
go func() {
|
||||
buf := make([]byte, 32)
|
||||
for i := 0; i < 2*iterations; i++ {
|
||||
for j := range buf {
|
||||
buf[j] = byte(i + j)
|
||||
}
|
||||
if _, err := ow.Write(buf); err != nil {
|
||||
errs <- err
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
wg.Wait()
|
||||
select {
|
||||
case err := <-errs:
|
||||
t.Fatalf("concurrent reader failed: %v", err)
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func firstBytes(p []tio.Packet) []byte {
|
||||
if len(p) == 0 {
|
||||
return nil
|
||||
}
|
||||
return p[0].Bytes
|
||||
}
|
||||
Reference in New Issue
Block a user