From 3b1004588d7d1a2ddd304922daaec838a46ef739 Mon Sep 17 00:00:00 2001 From: JackDoan Date: Mon, 3 Aug 2026 16:40:36 -0500 Subject: [PATCH] spicy offload chkpt --- firewall/packet.go | 9 ++ inside.go | 10 +- interface.go | 4 +- outside.go | 33 +++++-- outside_test.go | 95 ++++++++++++++++++- overlay/batch/batch.go | 11 ++- overlay/batch/coalesce_core.go | 102 ++++++++++++++------ overlay/batch/multi_coalesce.go | 116 +++++++++-------------- overlay/batch/multi_coalesce_test.go | 85 +++++++++++------ overlay/batch/passthrough.go | 6 +- overlay/batch/tcp_coalesce.go | 21 +++- overlay/batch/tcp_coalesce_bench_test.go | 13 ++- overlay/batch/udp_coalesce.go | 21 +++- 13 files changed, 363 insertions(+), 163 deletions(-) diff --git a/firewall/packet.go b/firewall/packet.go index 2cbfb5ea..8e2999e5 100644 --- a/firewall/packet.go +++ b/firewall/packet.go @@ -65,3 +65,12 @@ func (fp Packet) MarshalJSON() ([]byte, error) { "Fragment": fp.Fragment, }) } + +// ParsedPacket is a Packet plus the parse byproducts the RX path reuses +type ParsedPacket struct { + Packet + IPHdrLen int + // FragAny reports any fragmentation at all: MF flag or nonzero offset for IPv4, a fragment extension header for IPv6. + // Distinct from Packet.Fragment, which is true only for NON-FIRST fragments + FragAny bool +} diff --git a/inside.go b/inside.go index 32175b62..21a6b1e9 100644 --- a/inside.go +++ b/inside.go @@ -15,7 +15,7 @@ import ( "github.com/slackhq/nebula/routing" ) -func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packet, nb []byte, sendBatch *batch.SendBatch, rejectBuf []byte, q int, localCache firewall.ConntrackCache) { +func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.ParsedPacket, nb []byte, sendBatch *batch.SendBatch, rejectBuf []byte, q int, localCache firewall.ConntrackCache) { // borrowed: pkt.Bytes is owned by the originating tio.Queue and is // only valid until the next Read on that queue. Every consumer below // (parse, self-forward, handshake cache, sendInsideMessage) reads it @@ -74,7 +74,7 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe return } - hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) { + hostinfo, ready := f.getOrHandshakeConsiderRouting(&fwPacket.Packet, func(hh *HandshakeHostInfo) { // borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt // bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket), // so retaining segments past the loop is safe. @@ -105,7 +105,7 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe return } - dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache) + dropReason := f.firewall.Drop(fwPacket.Packet, false, hostinfo, f.pki.GetCAPool(), localCache) if dropReason == nil { f.sendInsideMessage(hostinfo, pkt, nb, sendBatch) } else { @@ -371,7 +371,7 @@ func (f *Interface) getOrHandshakeConsiderRouting(fwPacket *firewall.Packet, cac } func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) { - fp := &firewall.Packet{} + fp := &firewall.ParsedPacket{} err := newPacket(p, false, fp) if err != nil { f.l.Warn("error while parsing outgoing packet for firewall check", "error", err) @@ -379,7 +379,7 @@ func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubTyp } // check if packet is in outbound fw rules - dropReason := f.firewall.Drop(*fp, false, hostinfo, f.pki.GetCAPool(), nil) + dropReason := f.firewall.Drop(fp.Packet, false, hostinfo, f.pki.GetCAPool(), nil) if dropReason != nil { if f.l.Enabled(context.Background(), slog.LevelDebug) { f.l.Debug("dropping cached packet", diff --git a/interface.go b/interface.go index 89fc95e9..a6c512bf 100644 --- a/interface.go +++ b/interface.go @@ -360,7 +360,7 @@ func (f *Interface) listenOut(i int) { ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout) lhh := f.lightHouse.NewRequestHandler() h := &header.H{} - fwPacket := &firewall.Packet{} + fwPacket := &firewall.ParsedPacket{} nb := make([]byte, 12, 12) scratch := make([]byte, mtu) @@ -416,7 +416,7 @@ func (f *Interface) listenIn(queue tio.Queue, i int) { rejectBuf := make([]byte, mtu) arenaSize := batch.SendBatchCap * (udp.MTU + 32) sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, arenaSize) - fwPacket := &firewall.Packet{} + fwPacket := &firewall.ParsedPacket{} nb := make([]byte, 12, 12) conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout) diff --git a/outside.go b/outside.go index f6f0c6d1..378558d6 100644 --- a/outside.go +++ b/outside.go @@ -26,7 +26,7 @@ var ErrOutOfWindow = errors.New("out of window packet") // readOutsidePackets processes one received underlay packet. // Message payloads are decrypted IN PLACE, so packet must stay untouched // by the caller until the batcher for queue q has been flushed -func (f *Interface) readOutsidePackets(via ViaSender, scratch []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) { +func (f *Interface) readOutsidePackets(via ViaSender, scratch []byte, packet []byte, h *header.H, fwPacket *firewall.ParsedPacket, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) { err := h.Parse(packet) if err != nil { // Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors @@ -186,7 +186,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, scratch []byte, packet []b } } -func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, scratch []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) { +func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, scratch []byte, packet []byte, h *header.H, fwPacket *firewall.ParsedPacket, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) { // Successfully validated the thing. Get rid of the Relay header and the AEAD tag signedPayload := packet[header.Len : len(packet)-hostinfo.ConnectionState.dKey.Overhead()] // Pull the Roaming parts up here, and return in all call paths. @@ -315,7 +315,11 @@ var ( ) // newPacket validates and parses the interesting bits for the firewall out of the ip and sub protocol headers -func newPacket(data []byte, incoming bool, fp *firewall.Packet) error { +func newPacket(data []byte, incoming bool, fp *firewall.ParsedPacket) error { + // fp is reused across packets; reset the parse byproducts here so a + // parser's early-error return can't leak the previous packet's offsets. + fp.IPHdrLen = 0 + fp.FragAny = false if len(data) < 1 { return ErrPacketTooShort } @@ -330,7 +334,7 @@ func newPacket(data []byte, incoming bool, fp *firewall.Packet) error { return ErrUnknownIPVersion } -func parseV6(data []byte, incoming bool, fp *firewall.Packet) error { +func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error { dataLen := len(data) if dataLen < ipv6.HeaderLen { return ErrIPv6PacketTooShort @@ -356,6 +360,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error { switch proto { case layers.IPProtocolESP, layers.IPProtocolNoNextHeader: fp.Protocol = uint8(proto) + fp.IPHdrLen = offset fp.RemotePort = 0 fp.LocalPort = 0 fp.Fragment = false @@ -366,6 +371,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error { return ErrIPv6PacketTooShort } fp.Protocol = uint8(proto) + fp.IPHdrLen = offset fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6 icmptype := data[offset+1] switch icmptype { @@ -383,6 +389,9 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error { } fp.Protocol = uint8(proto) + // offset is the L4 header start: 40 for a plain packet, past the + // extension chain otherwise. The coalescer only accepts 40. + fp.IPHdrLen = offset if incoming { fp.RemotePort = binary.BigEndian.Uint16(data[offset : offset+2]) fp.LocalPort = binary.BigEndian.Uint16(data[offset+2 : offset+4]) @@ -400,6 +409,10 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error { return ErrIPv6PacketTooShort } + // Either way this packet is a fragment shape the coalescer must + // not touch, first fragment included. + fp.FragAny = true + // Check if this is the first fragment fragmentOffset := binary.BigEndian.Uint16(data[offset+2:offset+4]) &^ uint16(0x7) // Remove the reserved and M flag bits if fragmentOffset != 0 { @@ -441,7 +454,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error { return ErrIPv6CouldNotFindPayload } -func parseV4(data []byte, incoming bool, fp *firewall.Packet) error { +func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error { // Do we at least have an ipv4 header worth of data? if len(data) < ipv4.HeaderLen { return ErrIPv4PacketTooShort @@ -458,6 +471,10 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error { // Check if this is the second or further fragment of a fragmented packet. flagsfrags := binary.BigEndian.Uint16(data[6:8]) fp.Fragment = (flagsfrags & 0x1FFF) != 0 + // Any fragmentation at all (MF or offset): first fragments have readable + // ports for the firewall but must never be coalesced. + fp.FragAny = (flagsfrags & 0x3fff) != 0 + fp.IPHdrLen = ihl // Firewall handles protocol checks fp.Protocol = data[9] @@ -501,7 +518,7 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error { return nil } -func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, messageCounter uint64, out []byte, scratch []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) { +func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, messageCounter uint64, out []byte, scratch []byte, fwPacket *firewall.ParsedPacket, nb []byte, q int, localCache firewall.ConntrackCache) { err := newPacket(out, true, fwPacket) if err != nil { hostinfo.logger(f.l).Warn("Error while validating inbound packet", @@ -511,7 +528,7 @@ func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, messageCounte return } - dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache) + dropReason := f.firewall.Drop(fwPacket.Packet, true, hostinfo, f.pki.GetCAPool(), localCache) if dropReason != nil { f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, scratch, q) if f.l.Enabled(context.Background(), slog.LevelDebug) { @@ -523,7 +540,7 @@ func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, messageCounte return } - err = f.batchers[q].Commit(out, batch.SortKey{Epoch: hostinfo.ConnectionState.epoch, Counter: messageCounter}) + err = f.batchers[q].Commit(out, batch.SortKey{Epoch: hostinfo.ConnectionState.epoch, Counter: messageCounter}, fwPacket) if err != nil { f.l.Error("Failed to write to tun", "error", err) } diff --git a/outside_test.go b/outside_test.go index 4a24cae5..00060795 100644 --- a/outside_test.go +++ b/outside_test.go @@ -17,7 +17,7 @@ import ( ) func Test_newPacket(t *testing.T) { - p := &firewall.Packet{} + p := &firewall.ParsedPacket{} // length fails err := newPacket([]byte{}, true, p) @@ -96,7 +96,7 @@ func Test_newPacket(t *testing.T) { } func Test_newPacket_v6(t *testing.T) { - p := &firewall.Packet{} + p := &firewall.ParsedPacket{} // invalid ipv6 ip := layers.IPv6{ @@ -345,7 +345,7 @@ func Test_newPacket_v6(t *testing.T) { } func Test_newPacket_ipv6Fragment(t *testing.T) { - p := &firewall.Packet{} + p := &firewall.ParsedPacket{} ip := &layers.IPv6{ Version: 6, @@ -525,7 +525,7 @@ func BenchmarkParseV6(b *testing.B) { secondFrag = append(secondFrag, fragHeader...) secondFrag = append(secondFrag, []byte{0xde, 0xad, 0xbe, 0xef}...) - fp := &firewall.Packet{} + fp := &firewall.ParsedPacket{} b.Run("Normal", func(b *testing.B) { for i := 0; i < b.N; i++ { @@ -649,7 +649,7 @@ func serializeAH(ah *layers.IPSecAH) []byte { // host OS parses the real header, a firewall port/proto bypass. The fix makes parseV6 land // on the same offset the host does. func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) { - p := &firewall.Packet{} + p := &firewall.ParsedPacket{} const ( hdrLen = 40 // IPv6 header @@ -675,3 +675,88 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) { // the host delivers to, not the forged 443 at the overflowed offset. assert.Equal(t, uint16(22), p.LocalPort, "firewall must parse the real transport header, not the overflowed offset") } + +// Test_newPacket_parsedFields pins the ParsedPacket byproducts the RX +// batcher consumes: IPHdrLen (the true L4 offset) and FragAny (any fragment +// shape at all — unlike Packet.Fragment, which is port-oriented and true +// only for non-first fragments). +func Test_newPacket_parsedFields(t *testing.T) { + p := &firewall.ParsedPacket{} + + // Plain IPv4 TCP, IHL 20: L4 offset 20, no fragment shape. + v4 := make([]byte, 28) + v4[0] = 0x45 + v4[9] = firewall.ProtoTCP + binary.BigEndian.PutUint16(v4[6:8], 0x4000) // DF only + require.NoError(t, newPacket(v4, true, p)) + assert.Equal(t, 20, p.IPHdrLen) + assert.False(t, p.FragAny) + assert.False(t, p.Fragment) + + // IPv4 first fragment (MF set, offset 0): the firewall can read ports + // (Fragment false) but the coalescer must not touch it (FragAny true). + ff := make([]byte, 28) + ff[0] = 0x45 + ff[9] = firewall.ProtoUDP + binary.BigEndian.PutUint16(ff[6:8], 0x2000) // MF, offset 0 + require.NoError(t, newPacket(ff, true, p)) + assert.False(t, p.Fragment) + assert.True(t, p.FragAny) + assert.Equal(t, 20, p.IPHdrLen) + + // IPv4 non-first fragment (nonzero offset): both flags set. + nf := make([]byte, 28) + nf[0] = 0x45 + nf[9] = firewall.ProtoUDP + binary.BigEndian.PutUint16(nf[6:8], 0x00b9) + require.NoError(t, newPacket(nf, true, p)) + assert.True(t, p.Fragment) + assert.True(t, p.FragAny) + + // IPv4 with options (IHL 24): IPHdrLen tracks the real L4 offset. + opts := make([]byte, 32) + opts[0] = 0x46 + opts[9] = firewall.ProtoTCP + binary.BigEndian.PutUint16(opts[6:8], 0x4000) + require.NoError(t, newPacket(opts, true, p)) + assert.Equal(t, 24, p.IPHdrLen) + assert.False(t, p.FragAny) + + // Plain IPv6 TCP: L4 at 40. + v6 := make([]byte, 60) + v6[0] = 0x60 + v6[6] = firewall.ProtoTCP + require.NoError(t, newPacket(v6, true, p)) + assert.Equal(t, 40, p.IPHdrLen) + assert.False(t, p.FragAny) + + // IPv6 hop-by-hop then TCP: IPHdrLen lands past the extension header. + hbh := make([]byte, 60) + hbh[0] = 0x60 + hbh[6] = 0 // hop-by-hop + hbh[40] = firewall.ProtoTCP + hbh[41] = 0 // HdrExtLen 0 -> 8-byte header + require.NoError(t, newPacket(hbh, true, p)) + assert.Equal(t, 48, p.IPHdrLen) + assert.False(t, p.FragAny) + + // IPv6 first fragment: terminal proto resolved, FragAny set, Fragment not. + f6 := make([]byte, 60) + f6[0] = 0x60 + f6[6] = 44 // fragment extension header + f6[40] = firewall.ProtoUDP + require.NoError(t, newPacket(f6, true, p)) + assert.True(t, p.FragAny) + assert.False(t, p.Fragment) + assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol) + + // IPv6 non-first fragment: both set, walk stops at the fragment header. + f6n := make([]byte, 60) + f6n[0] = 0x60 + f6n[6] = 44 + f6n[40] = firewall.ProtoUDP + binary.BigEndian.PutUint16(f6n[42:44], 0x0008) + require.NoError(t, newPacket(f6n, true, p)) + assert.True(t, p.Fragment) + assert.True(t, p.FragAny) +} diff --git a/overlay/batch/batch.go b/overlay/batch/batch.go index eadb8573..81c21011 100644 --- a/overlay/batch/batch.go +++ b/overlay/batch/batch.go @@ -1,5 +1,7 @@ package batch +import "github.com/slackhq/nebula/firewall" + // SortKey identifies a packet's position in its sender's transmission order. // Epoch is a receiver-local ordinal for the tunnel (ConnectionState) that // decrypted the packet. A re-handshake replaces the tunnel outright — new @@ -15,9 +17,12 @@ type SortKey struct { type RxBatcher interface { // Commit stages pkt to be flushed by the batch. key must carry the - // packet's session epoch and message counter. The caller must keep pkt - // valid until the next Flush, and not re-use it. - Commit(pkt []byte, key SortKey) error + // packet's session epoch and message counter; pp must be the firewall's + // parse of this same packet. The caller must keep pkt valid until the + // next Flush, and not re-use it. pp, by contrast, is borrowed only for + // the duration of the call — the caller reuses one ParsedPacket per + // receive loop — so implementations must copy what they need from it. + Commit(pkt []byte, key SortKey, pp *firewall.ParsedPacket) error // Flush emits every staged packet. Packets are first sorted by key, so // within each protocol lane emission follows the sender's transmission // order regardless of arrival order. One shape may legally be overtaken diff --git a/overlay/batch/coalesce_core.go b/overlay/batch/coalesce_core.go index 1b1cf8d8..b0c00a81 100644 --- a/overlay/batch/coalesce_core.go +++ b/overlay/batch/coalesce_core.go @@ -40,34 +40,19 @@ type parsedIP struct { // On success, p.pkt is len-trimmed to the IP-declared length so callers // don't have to repeat the trim. wantProto is the IANA protocol number to // require (6 for TCP, 17 for UDP); ok=false for any other value. +// This is the standalone-lane-Commit entry; the dispatcher path uses +// parseIPAt, where the protocol was already resolved upstream. func parseIPPrologue(pkt []byte, wantProto byte) (parsedIP, bool) { var p parsedIP if len(pkt) < 20 { return p, false } - v := pkt[0] >> 4 - switch v { + switch pkt[0] >> 4 { case 4: - ihl := int(pkt[0]&0x0f) * 4 - if ihl != 20 { - return p, false - } if pkt[9] != wantProto { return p, false } - // Reject actual fragmentation (MF or non-zero frag offset). - if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 { - return p, false - } - totalLen := int(binary.BigEndian.Uint16(pkt[2:4])) - if totalLen > len(pkt) || totalLen < ihl { - return p, false - } - p.ipHdrLen = 20 - p.fk.isV6 = false - copy(p.fk.src[:4], pkt[12:16]) - copy(p.fk.dst[:4], pkt[16:20]) - p.pkt = pkt[:totalLen] + return parseIPv4Prologue(pkt) case 6: if len(pkt) < 40 { return p, false @@ -75,18 +60,77 @@ func parseIPPrologue(pkt []byte, wantProto byte) (parsedIP, bool) { if pkt[6] != wantProto { return p, false } - payloadLen := int(binary.BigEndian.Uint16(pkt[4:6])) - if 40+payloadLen > len(pkt) { - return p, false - } - p.ipHdrLen = 40 - p.fk.isV6 = true - copy(p.fk.src[:], pkt[8:24]) - copy(p.fk.dst[:], pkt[24:40]) - p.pkt = pkt[:40+payloadLen] - default: + return parseIPv6Prologue(pkt) + } + return p, false +} + +// parseIPAt is the dispatcher-path prologue: newPacket already resolved the +// L4 protocol and header offset once for the firewall, so the proto sniff is +// replaced by a cross-check of the caller's ipHdrLen. A plain header (v4: +// IHL 20, v6: exactly 40 — no options, no extension headers) is the only +// coalesceable shape, which is the same rule parseIPPrologue enforces +// through its own reads. +func parseIPAt(pkt []byte, ipHdrLen int) (parsedIP, bool) { + var p parsedIP + if len(pkt) < 20 { return p, false } + switch pkt[0] >> 4 { + case 4: + if ipHdrLen != 20 { + return p, false + } + return parseIPv4Prologue(pkt) + case 6: + if ipHdrLen != 40 || len(pkt) < 40 { + return p, false + } + return parseIPv6Prologue(pkt) + } + return p, false +} + +// parseIPv4Prologue is the shared IPv4 tail of the two prologue entries. +// The caller has verified len(pkt) >= 20 and either the protocol +// (parseIPPrologue) or the upstream-resolved header length (parseIPAt). +func parseIPv4Prologue(pkt []byte) (parsedIP, bool) { + var p parsedIP + ihl := int(pkt[0]&0x0f) * 4 + if ihl != 20 { + return p, false + } + // Reject actual fragmentation (MF or non-zero frag offset). On the + // dispatcher path FragAny was already gated; kept as defense in depth — + // a fragment folded into a superpacket would corrupt reassembly. + if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 { + return p, false + } + totalLen := int(binary.BigEndian.Uint16(pkt[2:4])) + if totalLen > len(pkt) || totalLen < ihl { + return p, false + } + p.ipHdrLen = 20 + p.fk.isV6 = false + copy(p.fk.src[:4], pkt[12:16]) + copy(p.fk.dst[:4], pkt[16:20]) + p.pkt = pkt[:totalLen] + return p, true +} + +// parseIPv6Prologue is the shared IPv6 tail; caller has verified +// len(pkt) >= 40 and version/proto-or-offset. +func parseIPv6Prologue(pkt []byte) (parsedIP, bool) { + var p parsedIP + payloadLen := int(binary.BigEndian.Uint16(pkt[4:6])) + if 40+payloadLen > len(pkt) { + return p, false + } + p.ipHdrLen = 40 + p.fk.isV6 = true + copy(p.fk.src[:], pkt[8:24]) + copy(p.fk.dst[:], pkt[24:40]) + p.pkt = pkt[:40+payloadLen] return p, true } diff --git a/overlay/batch/multi_coalesce.go b/overlay/batch/multi_coalesce.go index 3ee4cbe2..26781b42 100644 --- a/overlay/batch/multi_coalesce.go +++ b/overlay/batch/multi_coalesce.go @@ -6,7 +6,7 @@ import ( "log/slog" "slices" - "github.com/slackhq/nebula/iputil" + "github.com/slackhq/nebula/firewall" ) // MultiCoalescer stages plaintext packets with their (epoch, counter) sort @@ -47,9 +47,15 @@ type MultiCoalescer struct { staged []stagedPacket } +// stagedPacket also carries the scalars dispatch needs from the firewall's +// ParsedPacket: pp itself is reused by the caller per packet and must not be +// retained past Commit, so the relevant fields are copied by value here. type stagedPacket struct { - pkt []byte - key SortKey + pkt []byte + key SortKey + proto byte + fragAny bool + ipHdrLen uint16 } // NewMultiCoalescer builds a multi-lane batcher over w, based on available @@ -65,32 +71,18 @@ func NewMultiCoalescer(w io.Writer, l *slog.Logger) RxBatcher { return m } -// IANA protocol numbers for the IPv6 extension headers -// iputil.IPv6FindUpperProtocol can step over. The set here must match what -// that walker walks: it is the hot path's cheap pre-guard, so the walk is -// only paid when it can actually make progress. -const ( - ipProtoHopByHop = 0 - ipProtoRouting = 43 - ipProtoFragment = 44 - ipProtoAH = 51 - ipProtoDestOpts = 60 -) - -// isIPv6ExtHeader reports whether nh is an extension header the terminal- -// protocol walk knows how to step over. -func isIPv6ExtHeader(nh byte) bool { - switch nh { - case ipProtoHopByHop, ipProtoRouting, ipProtoFragment, ipProtoAH, ipProtoDestOpts: - return true - } - return false -} - -// Commit stages pkt for the next Flush. All parsing and lane dispatch is -// deferred to Flush so it runs on packets already in transmission order. -func (m *MultiCoalescer) Commit(pkt []byte, key SortKey) error { - m.staged = append(m.staged, stagedPacket{pkt: pkt, key: key}) +// Commit stages pkt for the next Flush. All lane dispatch is deferred to +// Flush so it runs on packets already in transmission order. pp is the +// firewall's parse of pkt — the single source of truth for the packet's +// protocol and L4 offset — and is only borrowed for this call. +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 } @@ -114,59 +106,43 @@ func compareStaged(a, b stagedPacket) int { return 1 } -// dispatch routes one packet to the appropriate lane based on IP version + -// L4 proto. On the success path the IP/TCP-or-UDP parse happens here once -// and the parsed struct is handed to the lane via commitParsed so the lane -// doesn't re-walk the header. -func (m *MultiCoalescer) dispatch(pkt []byte) error { - if len(pkt) < 20 { - return m.pt.enqueue(pkt) - } - v := pkt[0] >> 4 - var proto byte - switch v { - case 4: - proto = pkt[9] - case 6: - if len(pkt) < 40 { - return m.pt.enqueue(pkt) - } - proto = pkt[6] - if isIPv6ExtHeader(proto) { - // Walk to the terminal protocol so the packet routes to its flow's protocol lane. - // This protects flow ordering. - proto, _, _ = iputil.IPv6FindUpperProtocol(pkt) - } - default: - return m.pt.enqueue(pkt) - } - switch proto { +// dispatch routes one staged packet to its lane. +// The protocol and L4 offset come from the firewall's parse of the same packet. +// Any shape a lane can't coalesce seals every open chain in its lane +func (m *MultiCoalescer) dispatch(sp stagedPacket) error { + switch sp.proto { case ipProtoTCP: if m.tcp != nil { - info, ok := parseTCPBase(pkt) - if !ok { - // Unsupported TCP shape (IP options, fragments, ...). Its flow - // key is unknowable, so seal every open chain: dispatch runs in - // transmission order, and sealing is what keeps later data from - // extending a chain that would emit ahead of this packet. + if sp.fragAny { m.tcp.sealAllOpen() - m.tcp.addVerbatim(pkt) + m.tcp.addVerbatim(sp.pkt) return nil } - return m.tcp.commitParsed(pkt, info) + info, ok := parseTCPAt(sp.pkt, int(sp.ipHdrLen)) + if !ok { + m.tcp.sealAllOpen() + m.tcp.addVerbatim(sp.pkt) + return nil + } + return m.tcp.commitParsed(sp.pkt, info) } case ipProtoUDP: if m.udp != nil { - info, ok := parseUDP(pkt) - if !ok { + if sp.fragAny { m.udp.sealAllOpen() - m.udp.addVerbatim(pkt) + m.udp.addVerbatim(sp.pkt) return nil } - return m.udp.commitParsed(pkt, info) + info, ok := parseUDPAt(sp.pkt, int(sp.ipHdrLen)) + if !ok { + m.udp.sealAllOpen() + m.udp.addVerbatim(sp.pkt) + return nil + } + return m.udp.commitParsed(sp.pkt, info) } } - return m.pt.enqueue(pkt) + return m.pt.enqueue(sp.pkt) } // Flush sorts the staged batch into transmission order, replays it into the @@ -178,7 +154,7 @@ func (m *MultiCoalescer) Flush() error { var errs []error for _, sp := range m.staged { - if err := m.dispatch(sp.pkt); err != nil { + if err := m.dispatch(sp); err != nil { errs = append(errs, err) } } diff --git a/overlay/batch/multi_coalesce_test.go b/overlay/batch/multi_coalesce_test.go index ea138d6c..416c24a0 100644 --- a/overlay/batch/multi_coalesce_test.go +++ b/overlay/batch/multi_coalesce_test.go @@ -6,6 +6,7 @@ import ( "io" "testing" + "github.com/slackhq/nebula/firewall" "github.com/slackhq/nebula/test" ) @@ -48,19 +49,19 @@ func TestMultiCoalescerRoutesByProto(t *testing.T) { icmp[3] = 28 icmp[9] = 1 - if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay), k.next()); err != nil { + 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()); err != nil { + 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()); err != nil { + 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()); err != nil { + 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()); err != nil { + if err := m.Commit(icmp, k.next(), testPP(icmp)); err != nil { t.Fatal(err) } if err := m.Flush(); err != nil { @@ -89,13 +90,13 @@ func TestMultiCoalescerRestoresTransmissionOrder(t *testing.T) { // 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}); err != nil { + 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}); err != nil { + if err := m.Commit(buildTCPv4(1000, tcpAck, pay), SortKey{Epoch: 1, Counter: 1}, testPP(buildTCPv4(1000, tcpAck, pay))); err != nil { t.Fatal(err) } - if err := m.Commit(buildTCPv4(2200, tcpAck, pay), SortKey{Epoch: 1, Counter: 2}); err != nil { + 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 { @@ -115,10 +116,10 @@ func TestMultiCoalescerRestoresTransmissionOrder(t *testing.T) { // 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}); err != 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}); err != nil { + 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 { @@ -144,16 +145,16 @@ func TestMultiCoalescerRestoresOrderAcrossFlows(t *testing.T) { // 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}); err != nil { + 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}); err != nil { + 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}); err != nil { + 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}); err != nil { + 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 { @@ -195,10 +196,10 @@ func TestMultiCoalescerEpochOrdersAcrossRehandshake(t *testing.T) { 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}); err != nil { + 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}); err != nil { + 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 { @@ -227,10 +228,10 @@ func TestMultiCoalescerNoUSOFallsThrough(t *testing.T) { t.Fatal("UDP lane must not come up without USO") } - if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800)), k.next()); err != nil { + 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()); err != nil { + 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 { @@ -261,7 +262,7 @@ func TestMultiCoalescerNoOffloadsStillSorts(t *testing.T) { } // 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)}); err != nil { + if err := m.Commit(pkts[i], SortKey{Epoch: 1, Counter: uint64(i + 1)}, testPP(pkts[i])); err != nil { t.Fatal(err) } } @@ -325,13 +326,13 @@ func TestMultiCoalescerIPv6FragmentStaysInLane(t *testing.T) { m := newTestMultiCoalescer(t, w) k := &keySeq{epoch: 1} - if err := m.Commit(buildUDPv6Fragment(2000, 53, make([]byte, 512)), k.next()); err != nil { + 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()); err != nil { + 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()); err != nil { + 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 { @@ -358,19 +359,19 @@ func TestMultiCoalescerFragmentSealsUDPChains(t *testing.T) { m := newTestMultiCoalescer(t, w) k := &keySeq{epoch: 1} - if err := m.Commit(buildUDPv6(2000, 53, make([]byte, 800)), k.next()); err != nil { + 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()); err != nil { + 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()); err != nil { + 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()); err != nil { + 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()); err != nil { + 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 { @@ -398,10 +399,10 @@ func TestMultiCoalescerNoTSOFallsThrough(t *testing.T) { } pay := make([]byte, 1200) - if err := m.Commit(buildTCPv4(1000, tcpAck, pay), k.next()); err != nil { + 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()); err != nil { + 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 { @@ -414,3 +415,29 @@ func TestMultiCoalescerNoTSOFallsThrough(t *testing.T) { t.Errorf("TCP must pass through as 2 plain writes, got %d", len(w.writes)) } } + +// testPP derives the ParsedPacket newPacket would produce for the packet +// shapes the tests build: plain v4/v6, v4 with options or fragment bits set, +// and the single-fragment-header v6 shape from buildUDPv6Fragment. Anything +// unrecognizable stays zero (proto 0 routes to the passthrough lane). +func testPP(pkt []byte) *firewall.ParsedPacket { + pp := &firewall.ParsedPacket{} + if len(pkt) < 20 { + return pp + } + switch pkt[0] >> 4 { + case 4: + pp.Protocol = pkt[9] + pp.IPHdrLen = int(pkt[0]&0x0f) * 4 + pp.FragAny = binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 + case 6: + pp.Protocol = pkt[6] + pp.IPHdrLen = 40 + if pp.Protocol == 44 { // fragment extension header + pp.Protocol = pkt[40] + pp.IPHdrLen = 48 + pp.FragAny = true + } + } + return pp +} diff --git a/overlay/batch/passthrough.go b/overlay/batch/passthrough.go index 9d1c9308..7bf8d5c0 100644 --- a/overlay/batch/passthrough.go +++ b/overlay/batch/passthrough.go @@ -2,6 +2,8 @@ package batch import ( "io" + + "github.com/slackhq/nebula/firewall" ) // Passthrough is a RxBatcher that doesn't batch anything, it just accumulates and then sends packets. @@ -17,9 +19,7 @@ func NewPassthrough(w io.Writer) *Passthrough { } } -// Commit ignores the sort key: a bare Passthrough (no MultiCoalescer in -// front) emits in arrival order, exactly as before keys existed. -func (p *Passthrough) Commit(pkt []byte, _ SortKey) error { +func (p *Passthrough) Commit(pkt []byte, _ SortKey, _ *firewall.ParsedPacket) error { return p.enqueue(pkt) } diff --git a/overlay/batch/tcp_coalesce.go b/overlay/batch/tcp_coalesce.go index 165a8a4e..d0fc756b 100644 --- a/overlay/batch/tcp_coalesce.go +++ b/overlay/batch/tcp_coalesce.go @@ -117,12 +117,27 @@ type parsedTCP struct { // regardless of whether it's admissible for coalescing. Returns ok=false for non-TCP or malformed input. // Accepts IPv4 (no options or fragmentation) and IPv6 (no extension headers). func parseTCPBase(pkt []byte) (parsedTCP, bool) { - var p parsedTCP ip, ok := parseIPPrologue(pkt, ipProtoTCP) if !ok { - return p, false + return parsedTCP{}, false } - pkt = ip.pkt + return parseTCPTail(ip) +} + +// parseTCPAt is parseTCPBase for the dispatcher path: the packet is already +// known to be TCP and ipHdrLen is the upstream-resolved L4 offset (see parseIPAt). +func parseTCPAt(pkt []byte, ipHdrLen int) (parsedTCP, bool) { + ip, ok := parseIPAt(pkt, ipHdrLen) + if !ok { + return parsedTCP{}, false + } + return parseTCPTail(ip) +} + +// parseTCPTail layers the TCP-header parse on a validated IP prologue. +func parseTCPTail(ip parsedIP) (parsedTCP, bool) { + var p parsedTCP + pkt := ip.pkt p.fk = ip.fk p.ipHdrLen = ip.ipHdrLen diff --git a/overlay/batch/tcp_coalesce_bench_test.go b/overlay/batch/tcp_coalesce_bench_test.go index db73c554..29caaa79 100644 --- a/overlay/batch/tcp_coalesce_bench_test.go +++ b/overlay/batch/tcp_coalesce_bench_test.go @@ -4,6 +4,7 @@ import ( "encoding/binary" "testing" + "github.com/slackhq/nebula/firewall" "github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/test" ) @@ -169,16 +170,22 @@ func BenchmarkCommitNonCoalesceableTCP(b *testing.B) { // 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. +// 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++ { - pkt := pkts[i%len(pkts)] - if err := m.Commit(pkt, SortKey{Epoch: 1, Counter: uint64(i + 1)}); err != nil { + 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 { diff --git a/overlay/batch/udp_coalesce.go b/overlay/batch/udp_coalesce.go index c20466bd..1e23cc74 100644 --- a/overlay/batch/udp_coalesce.go +++ b/overlay/batch/udp_coalesce.go @@ -87,12 +87,27 @@ type parsedUDP struct { // Returns ok=false for non-UDP, malformed, or unsupported header shapes // (IPv4 with options/fragmentation, IPv6 with extension headers). func parseUDP(pkt []byte) (parsedUDP, bool) { - var p parsedUDP ip, ok := parseIPPrologue(pkt, ipProtoUDP) if !ok { - return p, false + return parsedUDP{}, false } - pkt = ip.pkt + return parseUDPTail(ip) +} + +// parseUDPAt is parseUDP for the dispatcher path: the packet is already +// known to be UDP and ipHdrLen is the upstream-resolved L4 offset (see parseIPAt). +func parseUDPAt(pkt []byte, ipHdrLen int) (parsedUDP, bool) { + ip, ok := parseIPAt(pkt, ipHdrLen) + if !ok { + return parsedUDP{}, false + } + return parseUDPTail(ip) +} + +// parseUDPTail layers the UDP-header parse on a validated IP prologue. +func parseUDPTail(ip parsedIP) (parsedUDP, bool) { + var p parsedUDP + pkt := ip.pkt p.fk = ip.fk p.ipHdrLen = ip.ipHdrLen