Compare commits

..

133 Commits

Author SHA1 Message Date
JackDoan 43fb7bff60 mutate packet 0, but, avoids a 100-byte copy 2026-08-04 14:59:58 -05:00
JackDoan ff672a3a1f remove dead expensive logging code 2026-08-04 13:40:07 -05:00
JackDoan 03269ffaa1 readOutsidePackets had too many args 2026-08-04 12:48:26 -05:00
JackDoan e196a7a7ca Revert "fable wants to DIY a hash"
This reverts commit e1d96b932a.
2026-08-04 11:54:53 -05:00
JackDoan d14b97977f Hostmap.QueryIndex is really hot 2026-08-04 11:34:02 -05:00
JackDoan e1d96b932a fable wants to DIY a hash 2026-08-04 09:56:39 -05:00
JackDoan d8d5ce344d stuff 2026-08-04 09:27:51 -05:00
JackDoan a3eef407b2 stuff 2026-08-04 09:03:40 -05:00
JackDoan 3b1004588d spicy offload chkpt 2026-08-03 16:40:36 -05:00
JackDoan 4cd433309b simplify 2026-08-03 15:48:57 -05:00
JackDoan 5ea48c1677 stop trying to interpret TCP, reorder via message counter and hostinfo-creation-order 2026-08-03 14:42:25 -05:00
JackDoan cdfba18ea5 crazy core pinning junk 2026-08-03 10:51:59 -05:00
JackDoan e04473f30b crazy core pinning junk 2026-08-03 10:51:37 -05:00
JackDoan 42c937d86f snip 2026-07-31 14:31:42 -05:00
JackDoan adfacc43c3 snip 2026-07-31 14:23:50 -05:00
JackDoan 096a06238a improve naming in batch 2026-07-31 14:03:30 -05:00
JackDoan bc70bf47f8 cap RX buffers in UDP rather than callers 2026-07-31 13:58:52 -05:00
JackDoan d7bcfb5d6b drop TxBatcher interface 2026-07-31 13:58:33 -05:00
JackDoan ddb90ad4b7 drop ECN support for this release 2026-07-31 13:42:15 -05:00
JackDoan 549de9fd29 drop ECN support for this release 2026-07-31 13:40:34 -05:00
JackDoan 93946faf7a drop ECN support for this release 2026-07-31 13:38:16 -05:00
JackDoan d3779b6a39 ecn: CE-mark on decap when the receive queue runs deep (nebula-as-AQM)
The tunnel's real bottleneck queue - the UDP receive buffer feeding the
decrypt loop - is invisible to every kernel AQM, so under overload it
regulates ECN-capable flows with tail-drop loss like it's 1993. Sample
SK_MEMINFO once per recvmmsg batch (tunnels.ecn_mark_threshold, fraction
of rcvbuf, 0=off) and treat depth beyond the threshold as an outer CE:
the existing RFC 6040 fold then CE-marks ECT inner packets and senders
back off without loss.
2026-07-31 13:09:50 -05:00
JackDoan 6783c90e72 ecn default disable eventually 2026-07-31 13:04:26 -05:00
JackDoan 245eb61444 don't pin when only one routine 2026-07-31 12:45:36 -05:00
JackDoan 9d8e830e4e disallow listen.batch < 1 2026-07-31 12:38:53 -05:00
JackDoan da8a640a56 QueueSet.Add errors if QueueSet is closed 2026-07-31 12:18:12 -05:00
JackDoan b7bf32240b plumb ECN through via relays 2026-07-31 12:14:56 -05:00
JackDoan b3002c2d13 clean out junk 2026-07-31 10:55:24 -05:00
JackDoan 575b97904d tun_linux: use IFF_TUN_EXCL to prevent multiple nebulas 2026-07-31 10:45:34 -05:00
JackDoan 16878eec1c fable fixes 2026-07-30 17:32:00 -05:00
JackDoan e8d6be1dd9 udp: resume partial sendmmsg in place instead of repacking
A partial success left the remaining entries' iovecs, sockaddrs, and
cmsgs fully intact, then threw them away and replanned the remainder
from bufs -- doubling the packing work exactly when the socket is
congested. Give sendFn a start offset so the drain resumes the same
prepared array at the first unsent entry, and skip a kernel-rejected
entry in place the same way. Only the GSO-disable path still replans,
since its entries change shape; it now rewinds precisely to the failed
run instead of the whole chunk, so entries already sent are never
duplicated.

New scripted tests pin the two paths that didn't exist before: a mid-
chunk rejected entry (drop it, resume the rest, start offsets advance)
and a mid-chunk EIO (GSO off, replay only the failed run, no dup of
already-sent packets).

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014ugV2edVqoz3tBvq9J6yWp
2026-07-29 17:55:13 -05:00
JackDoan d13c53db43 docs: mark six deferred review findings with TODOs
Behavior untouched; each marker records a known gap and the intended
fix so the next visit doesn't rediscover it: transient zero-sent
sendmmsg errors drop a whole run; the non-vnet Poll queue lacks the
post-wake drain loop; recvmmsg controllen resets touch every entry;
cached handshake packets flush one syscall each; the routines clamp in
activate() would blackhole surplus REUSEPORT sockets if it ever became
reachable; darwin WriteBatch burst-drops on EWOULDBLOCK.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014ugV2edVqoz3tBvq9J6yWp
2026-07-29 17:49:52 -05:00
JackDoan 7ee5e29758 docs: align comments with the code they describe
- multi_coalesce/batch: the ordering contract now states what Flush
  actually guarantees -- per-flow DATA order -- and names the two
  shapes later data may legally overtake (pure ACKs by design, and
  unparseable in-flow shapes as an accepted tradeoff).
- validVnetHdr claimed DATA_VALID makes the stack skip L4 checksum
  verification; the tun write path ignores that bit entirely. What the
  header buys is the absence of NEEDS_CSUM.
- tun_darwin Write said "only valid for single threaded use"; it is
  concurrency-safe and concurrent callers exist.
- udp_coalesce eviction comment said "Seal it" but never sets sealed.
- recordCapability: note the gauges are process-global while the state
  is per-socket (last writer wins).
- drop a stale tunReadBufSize reference.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014ugV2edVqoz3tBvq9J6yWp
2026-07-29 17:48:47 -05:00
JackDoan 9c68c60ba6 overlay/batch: don't seal the open slot on a pure ACK
Every non-coalesceable in-flow packet evicted the flow's open slot, so
a bidirectional connection's inbound data run was broken by each peer
ACK interleaved into it, largely defeating coalescing on concurrent
upload+download. A bare acknowledgment (zero payload, nothing beyond
ACK|PSH|ECE) carries no ordering obligation toward the flow's data --
delivered late it is just a stale ACK the receiver ignores -- so it
can ride the lane as a passthrough without the evict, same as kernel
GRO, which doesn't flush held data on pure ACKs. SYN/FIN/RST/CWR keep
sealing.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014ugV2edVqoz3tBvq9J6yWp
2026-07-29 17:39:13 -05:00
JackDoan 2190900107 overlay/batch: ship never-grown slots as plain writes
A slot that stays single-segment is byte-identical to the packet it
was seeded from, but flushSlot re-emitted it via WriteGSO with a
seeded pseudo-sum, forcing the kernel to software-checksum up to
~1400B that arrived with a perfectly valid checksum. Keep the borrowed
seed packet on the slot (valid until Flush per the Commit contract)
and emit it through the plain DATA_VALID path when numSeg is still 1
at flush time. appendPayload and mergeSlots only touch hdrBuf once
numSeg >= 2, so the raw bytes are pristine whenever the fast path
fires. This is every non-coalesced TCP/UDP packet: request/response
flows, many-flow fan-in, and each run's leftover tail.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014ugV2edVqoz3tBvq9J6yWp
2026-07-29 17:31:29 -05:00
JackDoan 69cf816f80 inside: mark traffic-out once per superpacket, not per segment
sendInsideEncrypt ran connectionManager.Out for every segment -- up to
~45 extra atomic stores per TSO superpacket, all inside writeLock when
boring crypto serializes encryption. One mark in sendInsideMessage
covers the whole superpacket on both the direct and relay paths.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014ugV2edVqoz3tBvq9J6yWp
2026-07-29 17:19:32 -05:00
JackDoan 8044b40c82 udp: degrade outer-ECN RX per-family on dual-stack sockets
A failed IPV6_RECVTCLASS probe disabled ECN RX entirely, taking down
working IPv4 (v4-mapped) delivery with it, while the opposite failure
already only degraded. Treat both directions the same: each family
degrades to Not-ECT independently, and only a full-family failure
turns the cmsg parsing off.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014ugV2edVqoz3tBvq9J6yWp
2026-07-29 17:17:56 -05:00
JackDoan fbcdd83359 overlay/tio: log dropped tun reads with bad virtio headers
decodeRead failures were silently swallowed; a kernel emitting an
unnegotiated GSO type would blackhole all tun traffic with nothing in
the logs. Debug-gated per the usual idiom so the happy path pays
nothing, which means plumbing the logger down through the offload
queueset.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014ugV2edVqoz3tBvq9J6yWp
2026-07-29 17:16:47 -05:00
JackDoan e9870686c8 overlay/tio: error on multi-segment WriteGSO with a bogus IP version
gsoTypeFromProto returns GSO_NONE when the IP version nibble is
neither 4 nor 6, so a multi-fragment superpacket went out as one
silent jumbo GSO_NONE packet -- exactly the silent mis-emission the
geometry checks promise not to allow.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014ugV2edVqoz3tBvq9J6yWp
2026-07-29 17:14:04 -05:00
JackDoan 9ce0596f5c tun: don't leak the QueueSet shutdown eventfd when Add fails
Both newTunGeneric error branches closed only the tun fd; the
Add-failure branch left the freshly created QueueSet (and its
shutdown eventfd) orphaned.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014ugV2edVqoz3tBvq9J6yWp
2026-07-29 17:12:46 -05:00
JackDoan 7455ac56eb outside: drop oversized test requests regardless of log level
The return lived inside a log-level-gated else-if, so with debug
logging off, control fell out of both switches. Nothing follows the
switch today, making it a silent drop by luck; any future code added
after the switch would have run for oversized Test requests only when
debug logging was disabled. Make the drop a guard clause.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014ugV2edVqoz3tBvq9J6yWp
2026-07-29 17:11:58 -05:00
JackDoan 0b0e45582a small bugs 2026-07-29 14:33:20 -05:00
JackDoan d43bb81ae2 batch stuff 2026-07-29 14:12:16 -05:00
JackDoan 8c06802f5f UDP stuff 2026-07-29 13:19:25 -05:00
JackDoan 5a8014585d silly e2e 2026-07-29 12:07:41 -05:00
JackDoan ffab005f9d dead code 2026-07-29 12:02:05 -05:00
JackDoan 38f11f6e3a udp: drop dead cmsg write 2026-07-29 11:59:29 -05:00
JackDoan 095a421708 silly opinionated tweak 2026-07-29 11:54:34 -05:00
JackDoan b39bae57ec tio.Offload.WriteGSO: reject 0-len packets, check seg lengths 2026-07-29 11:51:12 -05:00
JackDoan 5c0f6e2b5f unslop, improve the tio interface 2026-07-28 16:20:14 -05:00
JackDoan df8955177e unslop a bit 2026-07-28 15:00:33 -05:00
JackDoan 5c2a5607e5 unslop a bit 2026-07-28 14:56:51 -05:00
JackDoan ed88422770 unslop a bit 2026-07-28 14:42:04 -05:00
JackDoan 0b817c50b3 rework tun-side segmenentation checksums 2026-07-28 14:10:59 -05:00
JackDoan db54e05bfe checkpt 2026-07-28 12:17:55 -05:00
JackDoan 46a02b663e virtio: reject the GSO_ECN qualifier on non-TCP GSO types
35596c7 added the udp-l4-ecn-rejected test but only the ECN mask, so
UDP_L4|ECN validated as plain UDP_L4. Mirror virtio_net_hdr_to_skb and
refuse ECN on anything but TCPV4/TCPV6.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-28 10:54:03 -05:00
JackDoan b7d06615c0 Revert "tun: default pin_threads off"
This reverts commit fcf59343bd66dd3c879c8293c1429e95d9fd3ee2.
2026-07-27 16:23:34 -05:00
JackDoan 3803146bc9 improve WriteGSO again 2026-07-27 16:23:34 -05:00
JackDoan f9eb86df9d udp: smoke-test that GSO actually engages, on real sockets in CI
Nothing previously asserted the offload path works end to end -- a
silent fallback to per-packet sends would pass every test and only show
up as a throughput regression. Send a batch through a real StdConn over
loopback and assert (a) the run left as a single sendmmsg entry (GSO
engaged, via a spy around the real syscall) and (b) the kernel carved
the superpacket back into the exact original datagrams at the receiver.
Fails rather than skips when the probe reports no GSO on a UDP_SEGMENT-
capable kernel, so CI (make test, ubuntu-latest) guards engagement.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-27 16:23:34 -05:00
JackDoan 66cb98e13a udp: test the partial-sendmmsg rewind
The rewind (resume after the kernel accepts fewer entries than
submitted) was the hairiest untested logic in the write path; a bug
there silently duplicates or loses packets under backpressure. Give
batchWriter an injectable sendFn and drive WriteBatch through scripted
partial-acceptance sequences over a mixed GSO-run/plain batch, decoding
what "reached the wire" straight from the prepared iovecs rather than
the entryEnd bookkeeping under test. Also pins the zero-progress abort
and the EIO runtime GSO-disable replay.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-27 16:23:34 -05:00
JackDoan 3fe2cb970e udp: extract and test the GRO RX splitting 2026-07-27 16:23:34 -05:00
JackDoan 865dc9725c overlay/checksum: test each arch implementation directly
The correctness sweeps only exercised the public Checksum dispatcher, so
wherever it resolved to the gvisor fallback (non-AVX2 amd64, fallback
architectures) the suite compared gvisor against itself and the AVX2
assembly went untested -- silently green. Per-arch export_test.go files
now enumerate the hand-written implementations and every sweep runs
against the dispatcher plus each of them, skipping with an explicit
message when the running CPU can't execute one.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-27 16:21:10 -05:00
JackDoan 8cebecc087 overlay: test the checksum-seeding math against an RFC 1071 reference
pseudoSumIPv4/IPv6, foldOnceNoInvert, ipv4HdrChecksum (batch) and
foldComplement (tio/virtio) feed the virtio NEEDS_CSUM contract; a wrong
seed means every coalesced packet is silently dropped by the receiver
with nothing failing on our side. Check them against an independent
reference built from explicit RFC pseudo-header bytes -- deliberately
not the production checksum code -- including the carry/fold edge cases
and an end-to-end seed -> kernel-completion -> receiver-accepts
property.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-27 16:21:10 -05:00
JackDoan fb0d20f123 docs: note the GRO receive-scratch memory cost of listen.batch
The TX arena (128 x 9033B per routine) and GRO receive scratch
(listen.batch x 64KiB per socket) stay at their worst-case bounds by
design: the arena never grows past real demand and the GRO slots cannot
be smaller without truncating coalesced superpackets. Document the
listen.batch knob's memory implication so constrained hosts know what
to tune.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-27 16:21:10 -05:00
JackDoan 921ed4360a overlay/tio: validate WriteGSO geometry instead of silently dropping
The length checks were fishy on four counts: an empty hdr/transportHdr
with real payload returned nil (silent drop with a success signal); the
HdrLen/GSOSize/CsumStart uint16 conversions could wrap unchecked;
nothing verified transportHdr covers csum_start+csum_offset, so the
kernel's NEEDS_CSUM write could land in payload bytes; and there was no
total-size bound even though every length field involved is 16-bit.

Malformed geometry is now a real error, and a single 65535 total-length
guard makes all the u16 conversions exact.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-27 16:21:10 -05:00
JackDoan 6fb10c1c1a udp: make parseRecvCmsg's length check overflow-safe 2026-07-27 16:21:10 -05:00
JackDoan 8006b58758 util: unlock the OS thread when CPU pinning fails
PinThreadToCPU left the goroutine locked to its OS thread even when
sched_setaffinity failed. The lock only exists to make the affinity
stick; without it the kernel migrates the thread anyway, so a failed pin
kept a dedicated thread for zero benefit. Unwind on the error path.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-27 16:20:10 -05:00
JackDoan 56d2a3841d tun: default pin_threads off
Thread pinning trades scheduler freedom for TX-ring ordering; that's the
right trade on dedicated forwarders but not as a surprise default on
hosts sharing cores with other workloads. Make it opt-in and document
the default in the example config.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-27 16:20:10 -05:00
JackDoan 6c0305c7ae overlay/tio: derive WriteGSO geometry from non-empty fragments
GSOSize came from len(pays[0]) while the iovec build skips empty
fragments, so a leading empty fragment emitted a TSO/USO header with
gso_size == 0 -- virtio_net_hdr_to_skb rejects that with EINVAL and the
whole superpacket is lost. Compute gso_size from the first non-empty
fragment and use the non-empty count to decide superpacket vs plain.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-27 16:20:10 -05:00
JackDoan 5c0a5ee4be overlay/tio: guard Offload.Write against zero-length buffers
Write took &buf[0] before calling writeWithScratch, so the len==0 guard
in the helper could never run -- a zero-length buffer panicked on the
index instead of returning. Hoist the guard above the indexing and fold
writeWithScratch into Write since it was the only caller and duplicated
the iovec setup.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-27 16:20:10 -05:00
JackDoan 35596c7708 another VIRTIO_NET_HDR_GSO_ECN mistake 2026-07-27 16:20:10 -05:00
JackDoan d74c5ac5c5 overlay/batch: don't false-set PSH when merging a short-tail-sealed slot, clarify PSH vs sealing 2026-07-27 16:20:10 -05:00
JackDoan c06bfb46be overlay/batch: remove end-of-batch debug scaffolding
The Warn("==== end of batch ====") delimiter (and the `logged` flag that
fed it) was left over from debugging the cross-slot gap logging.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-27 16:14:33 -05:00
JackDoan be35d4f059 udp: downgrade the per-batch sendmmsg failure log to debug 2026-07-27 16:14:33 -05:00
JackDoan eea3c81626 udp: disable GSO at runtime when the kernel rejects a GSO send 2026-07-27 16:12:21 -05:00
Nate Brown 88872a8433 Don't fail on the batch at the first error (#1826) 2026-07-27 14:43:30 -05:00
Nate Brown 6bf424f749 Transmit a computed-zero UDP checksum as all ones (#1823)
yum yum
2026-07-24 19:25:49 -05:00
JackDoan 9688d32f5b make it nicer 2026-07-24 16:44:51 -05:00
JackDoan 8c91fa2699 fix it! 2026-07-24 16:44:51 -05:00
JackDoan 92d51c042e fix it! 2026-07-24 16:44:51 -05:00
JackDoan 0a0b2404a2 put locks around the replay window 2026-07-24 16:44:51 -05:00
JackDoan ef0e3015f9 decrypt in place 2026-07-24 16:44:05 -05:00
JackDoan c6ebe71c08 silly optimization 2026-07-24 16:39:03 -05:00
JackDoan 1d768ac4e4 tio: accept VIRTIO_NET_HDR_GSO_ECN-qualified superpackets
TUN_F_TSO_ECN is negotiated, so once ECN feedback flows the kernel hands
us TSO superpackets typed TCPV4|GSO_ECN (CWR set). protoFromGSOType
treated the qualifier bit as an unknown type and the read path dropped
every such superpacket - a latent bug that only fires when a congested
hop CE-marks the flow, exactly when drops hurt most. The segmenter
already handles CWR (first segment only); just mask the bit.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-24 16:39:03 -05:00
JackDoan 69fa9e4a2e unslop some comments 2026-07-24 16:39:03 -05:00
JackDoan ed55cf40d5 batch: back SendBatch with Arena instead of a hand-rolled slab
SendBatch.Reserve duplicated Arena's grow-on-demand logic byte for byte.
Use an Arena for the slot backing so the borrow/grow/recycle semantics
live in one place.
2026-07-24 16:39:03 -05:00
JackDoan a6ae44ddb1 batch: move shared-arena Reset ownership from lanes to their owner 2026-07-24 16:39:03 -05:00
JackDoan 7cc37323a8 re-align to master 2026-07-24 16:39:03 -05:00
JackDoan 68e3fae870 some tests 2026-07-24 16:39:03 -05:00
JackDoan 720990ddcd simplify making new Queues 2026-07-24 16:39:03 -05:00
JackDoan 05443523bd make service test less annoying 2026-07-24 16:39:03 -05:00
JackDoan 99bf613a2c checkpt 2026-07-24 16:39:03 -05:00
JackDoan 9e7c783eb3 checkpt 2026-07-24 16:39:03 -05:00
JackDoan 8b14a6ee56 more ram -> more speed 2026-07-24 16:39:03 -05:00
JackDoan ae17513bbf more fixes! 2026-07-24 16:39:03 -05:00
JackDoan 1aca2f75ae more fixes! 2026-07-24 16:39:03 -05:00
JackDoan c7918d1096 lint 2026-07-24 16:39:03 -05:00
JackDoan 44dd2e9ca4 datapath: fix 12 correctness findings from tun/UDP offload review
Multi-disciplinary correctness review of the batched tun / GSO-GRO / sendmmsg
rework. Each fix has a regression test; the merged tree builds on
linux/darwin/openbsd/windows/freebsd/netbsd, vets clean, passes the unit and
e2e suites, and is -race clean.

Critical:
- C1 zero-length inner UDP datagram no longer panics the process (remote DoS):
  the UDP coalescer routes payLen==0 to passthrough instead of seeding a GSO
  slot, and WriteGSO skips empty payload iovecs as defense in depth.
- C2 segmenter no longer corrupts inner headers when gsoSize < headerLen: the
  L3+L4 header is snapshotted once and each segment stamped from the copy,
  replacing the destructive overlapping in-place slide (SegmentTCP + SegmentUDP).

High:
- H1 applyOuterECN updates the IPv4 header checksum (RFC 1624 incremental) when
  folding outer CE into the inner ToS, so passthrough packets are no longer
  dropped by the peer stack.
- H2 the GRO reject path caps the borrowed RX segment ([:n:n]) so a reject can
  no longer overrun into the next coalesced segment's Nebula header. Note:
  oversized ICMPv6 rejects that need >16B beyond the segment are now refused
  rather than sent under GRO (safe; see TOFIX.md for the scratch-buffer follow-up).
- H3 WriteBatch falls back to per-packet WriteTo for a chunk when writeSockaddr
  fails, so one bad-family destination costs only its own packet, not the batch.
- H4 UserDevice.Readers returns N distinct queue wrappers with private buffers
  (sharing the pipes) so concurrent readers no longer race/overwrite borrowed
  packet bytes.
- H5 Poll.Close / Offload.Close no longer null t.fd (matching master's
  tunFile.Close), removing the data race with a concurrent readOne load.

Medium/Low:
- M1 the UDP GSO 127-segment gate moved from kernel >=5.5 to >=6.9 (the real
  UDP_MAX_SEGMENTS 64->128 threshold), avoiding EINVAL + per-packet fallback on
  5.5-6.8 kernels.
- M2 NewMultiQueueReader replays the offload mask newTun actually negotiated
  instead of the TSO-only mask, so adding a queue no longer disables USO
  device-wide; the advertised USO capability derives from the same mask.
- M3 the shutdown eventfd is closed in pollQueueSet.Close / offloadQueueSet.Close
  (double-close guarded), fixing the per-lifecycle fd leak.
- M4 dual-stack ECN selects the cmsg by address family, not socket family: RX
  parseRecvCmsg reads both IP_TOS and IPV6_TCLASS; TX writeEntryCmsg stamps
  IP_TOS for v4/v4-mapped dests and IPV6_TCLASS for v6 (on-host verified).
- L1 newPoll no longer closes the fd on failure (matching newOffload), removing
  the double-close on QueueSet.Add error.
2026-07-24 16:39:03 -05:00
JackDoan 733dc06192 make mobile happy 2026-07-24 16:39:03 -05:00
JackDoan 7fee3a97b2 correctly shutdown the pprofserver 2026-07-24 16:39:03 -05:00
JackDoan 1e218737dc SendVia: don't emit a zero-length packet when prepareSendVia fails 2026-07-24 16:39:03 -05:00
JackDoan 243c920f88 adapt Control lifecycle tests to the batched tio.Queue Device interface 2026-07-24 16:39:03 -05:00
JackDoan 9bdab873f2 udp setsockopt correctness fixes 2026-07-24 16:39:03 -05:00
JackDoan a081fba023 use less ram pls 2026-07-24 16:39:03 -05:00
JackDoan b50d6276e3 clean up a comment a bit 2026-07-24 16:39:03 -05:00
JackDoan 3b16f1adb6 drop in a logger 2026-07-24 16:39:03 -05:00
JackDoan 410bac9688 go mod tidy 2026-07-24 16:39:03 -05:00
JackDoan 22adae0b8c lint 2026-07-24 16:39:03 -05:00
JackDoan f7cc437d88 fix 2026-07-24 16:39:03 -05:00
JackDoan a66af843d1 faster
grr heap usage!
2026-07-24 16:38:32 -05:00
JackDoan 67a742ddfb no 2026-07-24 16:38:32 -05:00
JackDoan 5681e510c4 use clear() 2026-07-24 16:38:32 -05:00
JackDoan 9104dc4c34 remove udp-level RX reorder buf 2026-07-24 16:38:32 -05:00
JackDoan f31f6c5d1f make relays take the fast path maybe 2026-07-24 16:37:30 -05:00
JackDoan 264a25337b scoot pinning around 2026-07-24 16:37:30 -05:00
JackDoan 694414771c scoot stuff around for e2e 2026-07-24 16:37:30 -05:00
JackDoan ca84bcb38f disable sort-on-RX, CPU pinning seems to work for now 2026-07-24 16:37:30 -05:00
JackDoan f04dd3bcc3 switch to ASM vector checksum 2026-07-24 16:37:30 -05:00
JackDoan 187afac7b5 GSO/GRO offloads, with TCP+ECN and UDP support 2026-07-24 16:37:30 -05:00
JackDoan e7b121c82f better and batched tun interface 2026-07-24 16:33:52 -05:00
Nate Brown 72bf111209 Add an e2e Drop exit type and a roaming recovery measurement (#1819)
smoke-extra / freebsd-amd64 (push) Failing after 15s
smoke-extra / linux-amd64-ipv6disable (push) Failing after 15s
smoke-extra / netbsd-amd64 (push) Failing after 14s
smoke-extra / openbsd-amd64 (push) Failing after 15s
smoke-extra / linux-386 (push) Failing after 16s
smoke / Run multi node smoke test (push) Failing after 1m37s
Build and test / Static checks (push) Successful in 18s
Build and test / Test linux (push) Failing after 58s
Build and test / Test linux-boringcrypto (push) Failing after 2m45s
Build and test / Test linux-pkcs11 (push) Failing after 2m10s
Build and test / Cross-build linux-arm (push) Successful in 3m11s
Build and test / Cross-build linux-mips (push) Successful in 3m53s
Build and test / Cross-build linux-other (push) Successful in 3m16s
Build and test / Cross-build windows (push) Successful in 1m2s
Build and test / Cross-build freebsd (push) Successful in 1m36s
Build and test / Cross-build netbsd (push) Successful in 1m36s
Build and test / Cross-build openbsd (push) Successful in 1m37s
Build and test / Cross-build mobile (push) Successful in 3m23s
smoke-extra / Run windows smoke test (push) Has been cancelled
Build and test / Test macos (push) Has been cancelled
Build and test / Test windows (push) Has been cancelled
Build and test / CI status (push) Has been cancelled
2026-07-23 17:02:02 -05:00
Nate Brown 1617897043 v1.11.0 changelog (#1792)
smoke-extra / freebsd-amd64 (push) Failing after 25s
smoke-extra / linux-amd64-ipv6disable (push) Failing after 13s
smoke-extra / netbsd-amd64 (push) Failing after 13s
smoke-extra / openbsd-amd64 (push) Failing after 11s
smoke-extra / linux-386 (push) Failing after 12s
smoke / Run multi node smoke test (push) Failing after 1m39s
Build and test / Static checks (push) Successful in 2m9s
Build and test / Test linux (push) Failing after 1m3s
Build and test / Test linux-boringcrypto (push) Failing after 2m48s
Build and test / Test linux-pkcs11 (push) Failing after 2m0s
Build and test / Cross-build linux-arm (push) Successful in 3m17s
Build and test / Cross-build linux-mips (push) Successful in 4m5s
Build and test / Cross-build linux-other (push) Successful in 3m23s
Build and test / Cross-build windows (push) Successful in 1m4s
Build and test / Cross-build freebsd (push) Successful in 1m42s
Build and test / Cross-build netbsd (push) Successful in 1m38s
Build and test / Cross-build openbsd (push) Successful in 1m43s
Build and test / Cross-build mobile (push) Successful in 3m34s
smoke-extra / Run windows smoke test (push) Has been cancelled
Build and test / Test macos (push) Has been cancelled
Build and test / Test windows (push) Has been cancelled
Build and test / CI status (push) Has been cancelled
2026-07-23 13:13:45 -05:00
Nate Brown f8775bb6ca Use go 1.26 (latest 1.26.5) (#1818) 2026-07-23 10:36:20 -05:00
Nate Brown 15f0f0d5d0 Be less verbose with handshake send errors (#1810) 2026-07-23 09:26:45 -05:00
Nate Brown 7902ce674e Rebind for MacOS (#1816)
Co-authored-by: Jack Doan <me@jackdoan.com>
2026-07-23 09:26:24 -05:00
Nate Brown c2fbe215e6 Fix a test race, make dns server reload/restart safer (#1815)
smoke-extra / freebsd-amd64 (push) Failing after 54s
smoke-extra / linux-amd64-ipv6disable (push) Failing after 11s
smoke-extra / netbsd-amd64 (push) Failing after 11s
smoke-extra / openbsd-amd64 (push) Failing after 22s
smoke-extra / linux-386 (push) Failing after 23s
smoke / Run multi node smoke test (push) Failing after 1m26s
Build and test / Static checks (push) Successful in 2m11s
Build and test / Test linux (push) Failing after 1m7s
Build and test / Test linux-boringcrypto (push) Failing after 2m41s
Build and test / Test linux-pkcs11 (push) Failing after 2m2s
Build and test / Cross-build linux-arm (push) Successful in 3m5s
Build and test / Cross-build linux-mips (push) Successful in 3m48s
Build and test / Cross-build linux-other (push) Successful in 3m8s
Build and test / Cross-build windows (push) Successful in 1m4s
Build and test / Cross-build freebsd (push) Successful in 1m36s
Build and test / Cross-build netbsd (push) Successful in 1m33s
Build and test / Cross-build openbsd (push) Successful in 1m32s
Build and test / Cross-build mobile (push) Successful in 3m16s
smoke-extra / Run windows smoke test (push) Has been cancelled
Build and test / Test macos (push) Has been cancelled
Build and test / Test windows (push) Has been cancelled
Build and test / CI status (push) Has been cancelled
2026-07-22 15:05:10 -05:00
dependabot[bot] 94ac6db4ca Bump the golang-x-dependencies group across 1 directory with 5 updates (#1800)
Bumps the golang-x-dependencies group with 3 updates in the / directory: [golang.org/x/crypto](https://github.com/golang/crypto), [golang.org/x/net](https://github.com/golang/net) and [golang.org/x/sync](https://github.com/golang/sync).


Updates `golang.org/x/crypto` from 0.53.0 to 0.54.0
- [Commits](https://github.com/golang/crypto/compare/v0.53.0...v0.54.0)

Updates `golang.org/x/net` from 0.56.0 to 0.57.0
- [Commits](https://github.com/golang/net/compare/v0.56.0...v0.57.0)

Updates `golang.org/x/sync` from 0.21.0 to 0.22.0
- [Commits](https://github.com/golang/sync/compare/v0.21.0...v0.22.0)

Updates `golang.org/x/sys` from 0.46.0 to 0.47.0
- [Commits](https://github.com/golang/sys/compare/v0.46.0...v0.47.0)

Updates `golang.org/x/term` from 0.44.0 to 0.45.0
- [Commits](https://github.com/golang/term/compare/v0.44.0...v0.45.0)

---
updated-dependencies:
- dependency-name: golang.org/x/crypto
  dependency-version: 0.54.0
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: golang-x-dependencies
- dependency-name: golang.org/x/net
  dependency-version: 0.57.0
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: golang-x-dependencies
- dependency-name: golang.org/x/sync
  dependency-version: 0.22.0
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: golang-x-dependencies
- dependency-name: golang.org/x/sys
  dependency-version: 0.47.0
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: golang-x-dependencies
- dependency-name: golang.org/x/term
  dependency-version: 0.45.0
  dependency-type: direct:production
  update-type: version-update:semver-minor
  dependency-group: golang-x-dependencies
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-07-22 15:51:22 -04:00
dependabot[bot] a60350e34e Bump actions/setup-go from 6 to 7 (#1807)
Bumps [actions/setup-go](https://github.com/actions/setup-go) from 6 to 7.
- [Release notes](https://github.com/actions/setup-go/releases)
- [Commits](https://github.com/actions/setup-go/compare/v6...v7)

---
updated-dependencies:
- dependency-name: actions/setup-go
  dependency-version: '7'
  dependency-type: direct:production
  update-type: version-update:semver-major
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-07-22 15:49:45 -04:00
John Maguire 58f3b6fda7 Document rootless Nebula in example service script (#1814) 2026-07-21 19:32:33 -04:00
Jack Doan a99699e370 before removing a pending hostinfo in handshake_manager, make sure it's the one we wanted to delete (#1811) 2026-07-21 10:31:00 -05:00
Jack Doan 3615a79b8b add locks around replay window updates (#1802) 2026-07-20 10:28:53 -05:00
Nate Brown 147c202c27 Swap back to a blocking udp socket, test shutdown(2) (#1806)
Co-authored-by: Jack Doan <me@jackdoan.com>
2026-07-17 15:16:45 -05:00
John Maguire e290a6892f Fix relay re-establishment for handshake on Disestablised entry (#1805)
handleOutsideRelayPacket filled ViaSender.remoteIdx with relay.RemoteIndex,
an index from the relay peer's index space, but the rescue in
sendHandshakeResponse looks that value up in relayForByIdx, which is keyed
by local index. The lookup could never hit, so a terminal relay entry left
Disestablished by a one-sided teardown stayed Disestablished even after a
valid handshake arrived over it. The responder's first transmit then failed
to find an Established relay, deleted its only relay entry, and every
subsequent send was silently dropped until dead-tunnel detection forced a
re-handshake.
2026-07-17 11:57:47 -04:00
109 changed files with 8028 additions and 4209 deletions
+6 -6
View File
@@ -12,9 +12,9 @@ jobs:
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v7
- uses: actions/setup-go@v6 - uses: actions/setup-go@v7
with: with:
go-version: '1.25' go-version: '1.26'
check-latest: true check-latest: true
- name: Build - name: Build
@@ -38,9 +38,9 @@ jobs:
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v7
- uses: actions/setup-go@v6 - uses: actions/setup-go@v7
with: with:
go-version: '1.25' go-version: '1.26'
check-latest: true check-latest: true
- name: Build - name: Build
@@ -78,9 +78,9 @@ jobs:
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v7
- uses: actions/setup-go@v6 - uses: actions/setup-go@v7
with: with:
go-version: '1.25' go-version: '1.26'
check-latest: true check-latest: true
- name: Import certificates - name: Import certificates
+6 -6
View File
@@ -32,9 +32,9 @@ jobs:
- uses: actions/checkout@v7 - uses: actions/checkout@v7
- uses: actions/setup-go@v6 - uses: actions/setup-go@v7
with: with:
go-version: '1.25' go-version: '1.26'
check-latest: true check-latest: true
- name: add hashicorp source - name: add hashicorp source
@@ -64,9 +64,9 @@ jobs:
- uses: actions/checkout@v7 - uses: actions/checkout@v7
- uses: actions/setup-go@v6 - uses: actions/setup-go@v7
with: with:
go-version: '1.25' go-version: '1.26'
check-latest: true check-latest: true
- name: add hashicorp source - name: add hashicorp source
@@ -90,9 +90,9 @@ jobs:
- uses: actions/checkout@v7 - uses: actions/checkout@v7
- uses: actions/setup-go@v6 - uses: actions/setup-go@v7
with: with:
go-version: '1.25' go-version: '1.26'
check-latest: true check-latest: true
# WSL2 + Ubuntu so the smoke can run a real linux peer with its own # WSL2 + Ubuntu so the smoke can run a real linux peer with its own
+2 -2
View File
@@ -20,9 +20,9 @@ jobs:
- uses: actions/checkout@v7 - uses: actions/checkout@v7
- uses: actions/setup-go@v6 - uses: actions/setup-go@v7
with: with:
go-version: '1.25' go-version: '1.26'
check-latest: true check-latest: true
- name: build - name: build
+7 -7
View File
@@ -20,9 +20,9 @@ jobs:
- uses: actions/checkout@v7 - uses: actions/checkout@v7
- uses: actions/setup-go@v6 - uses: actions/setup-go@v7
with: with:
go-version: '1.25' go-version: '1.26'
check-latest: true check-latest: true
- name: Install goimports - name: Install goimports
@@ -42,7 +42,7 @@ jobs:
- name: golangci-lint - name: golangci-lint
uses: golangci/golangci-lint-action@v9 uses: golangci/golangci-lint-action@v9
with: with:
version: v2.5 version: v2.12
test: test:
name: Test ${{ matrix.name }} name: Test ${{ matrix.name }}
@@ -80,9 +80,9 @@ jobs:
- uses: actions/checkout@v7 - uses: actions/checkout@v7
- uses: actions/setup-go@v6 - uses: actions/setup-go@v7
with: with:
go-version: '1.25' go-version: '1.26'
check-latest: true check-latest: true
- name: Build - name: Build
@@ -125,9 +125,9 @@ jobs:
- uses: actions/checkout@v7 - uses: actions/checkout@v7
- uses: actions/setup-go@v6 - uses: actions/setup-go@v7
with: with:
go-version: '1.25' go-version: '1.26'
check-latest: true check-latest: true
- name: Build ${{ matrix.name }} - name: Build ${{ matrix.name }}
+82
View File
@@ -7,6 +7,88 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [Unreleased] ## [Unreleased]
## [1.11.0] - 2026-07-23
See the [v1.11.0](https://github.com/slackhq/nebula/milestone/25?closed=1) milestone for a complete list of changes.
### Breaking
- Logging has switched from logrus to Go's structured `slog`. Log output changes: levels are upper case
(`level=INFO`), trace prints as `level=DEBUG-4`, timestamps are always RFC3339Nano and `logging.timestamp_format`
is ignored, and some messages were reworded. Review any log parsing before upgrading. This is also an API break
for embedders, as constructors now take a `*slog.Logger`. (#1672, #1734, #1621)
- `firewall.inbound_action` and `firewall.outbound_action` (used to set reject vs. drop policy) were each being
applied to the opposite direction, that is now corrected. This only affects how blocked packets are answered, not
which packets the firewall allows or denies. If you set either of these you are getting the behavior of the other
one today and likely want to swap them before upgrading. (#1798)
- On Windows, Nebula now installs WFP PERMIT filters for the nebula adapter and the listener port by default. WFP
sits below Windows Defender Firewall, so any WDF inbound rules you rely on for either will no longer apply. Set
`tun.windows_bypass_wdf` and `listen.windows_bypass_wdf` to false to leave WDF in charge. (#1710)
- On Windows, the nebula device is now set to the `private` network category instead of whatever Windows decided,
which is usually `Public`. This makes the host firewall less restrictive on the overlay. Set
`tun.network_category` to `unset` to keep the old behavior. (#1710)
- Reject packets for non-TCP now use ICMP code 13, communication administratively prohibited, instead of code 3,
port unreachable. Anything keying off the old code needs updating. (#1766, #1768)
- The SSH debug server's profiling commands are now confined to `sshd.sandbox_dir`, which defaults to
`$TMP/nebula-debug`. Relative paths resolve inside it and absolute paths outside it are rejected, so anything
scripting `start-cpu-profile`, `save-heap-profile`, or `save-mutex-profile` with a path elsewhere needs the
directory set. The directory is not created for you. (#1622)
### Added
- Sign the Windows release binaries. (#1718)
- Generate IPv6 reject packets, matching the existing IPv4 behavior. (#1766, #1767, #1768)
- Accept `-` in `nebula-cert` to read from stdin or write to stdout. (#1714)
- Search for both `config.yml` and `config.yaml` in service and command line modes. (#1717)
- Add version labels to the Docker/OCI images. (#1772)
- Rebind the listener and re-query lighthouses on macOS when the underlay network changes, so devices moving
between wifi and wired or between networks recover without waiting for dead tunnel detection. Controlled by
`listen.rebind_on_network_change` (default `true`, not reloadable). (#1816)
### Changed
- Reload the firewall when the unsafe networks in the certificate change. (#1719)
- Reconfigure, start, and stop the stats listener on a config reload instead of requiring a restart. (#1670)
- Update a static host's addresses when they change on reload. (#1713)
- Don't require a port on ICMP firewall rules. (#1609)
- Connection track ICMP traffic. (#1602)
- Return `NODATA` instead of `NXDOMAIN` from the DNS server for a name that exists but has no record of the
requested type, so clients that query `AAAA` first (busybox/Alpine) fall through to `A`. (#1668)
- Record the local host's details in the DNS server. (#1716)
- Install Windows unsafe routes as link routes. (#1709)
- Reduce relay handshake log spam, and only log a handshake send error at error level when the remote list
changes. (#1733, #1765, #1810)
- Start, stop, and reload subsystems (DNS, stats, conntrack, ssh, punchy) cleanly without leaking goroutines. (#1640, #1654, #1661, #1667, #1669, #1708, #1806, #1815)
- `Control` is now safe to stop and wait on from any lifecycle state, and a new `Control.Wait` blocks until nebula
has fully stopped and returns the first fatal reader error. Failed starts release the udp sockets and tun fd
instead of leaking them. (#1794)
- Trigger an immediate lighthouse update when reconnecting to or adding a lighthouse instead of waiting for the next update tick. (#1645)
- Bring the Darwin and OpenBSD tun implementations in line with the other BSDs. (#1703)
- Update to build against go v1.26. (#1818)
- Various dependency updates. (#1586, #1587, #1604, #1617, #1618, #1627, #1628, #1629, #1652, #1664, #1665, #1697, #1721, #1732, #1742, #1743, #1750, #1763, #1771, #1782, #1800, #1807)
### Fixed
- Fix a data race on a host's remote address that could send packets to the wrong address during a roam. (#1773)
- Fix tunnels that could permanently escape connection manager monitoring. (#1752)
- Fix a crash when reloading the SSH server's trusted keys. (#1787)
- Fix hostmap corruption when a host has multiple overlay addresses. Each address now gets its own list instead of
a single shared chain, which also fixes two latent bugs on the add and makePrimary paths. (#1788, #1790)
- Apply `remote_allow_list` IPv4 rules to 4-in-6 mapped addresses. (#1786)
- Don't panic in the DNS server on a short or empty query name. (#1635)
- Advance the replay window on relayed packets so a relay drops replayed frames instead of re-forwarding them. (#1751)
- Fix a race in relay state handling. (#1753)
- Lock replay window updates so concurrent readers can't corrupt it. (#1802)
- Reject malformed handshakes more reliably, including invalid ed25519 key lengths. (#1601, #1756)
- Properly handle `closetunnel` packets. (#1638)
- Fix an IPv6 extension-header length overflow that could make the firewall parse the wrong protocol and ports. (#1789)
- Fix relay re-establishment when a handshake arrives over a relay entry that a one-sided teardown left
`Disestablished`, which silently dropped every send until dead tunnel detection forced a re-handshake. (#1805)
- Don't build new relay state on a tunnel that was just discarded. (#1796)
- Don't delete the wrong pending hostinfo in the handshake manager. (#1811)
- Don't call the packet reader after a UDP error on Darwin. (#1755)
- Open the FreeBSD tun device non blocking. (#1666)
## [1.10.3] - 2026-02-06 ## [1.10.3] - 2026-02-06
### Security ### Security
+96
View File
@@ -0,0 +1,96 @@
//go:build linux && !android && !e2e_testing
package main
import (
"fmt"
"net/netip"
"os"
"path/filepath"
"runtime"
"testing"
"time"
"github.com/slackhq/nebula"
"github.com/slackhq/nebula/cert"
cert_test "github.com/slackhq/nebula/cert_test"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/test"
"github.com/stretchr/testify/require"
)
// TestControlStopClosesOnTimer reproduces the dnclient lifecycle: nebula runs as
// a library, and on a config update dnclient calls Stop() in-process to tear the
// old instance down before starting a new one. This boots a real nebula (real
// blocking UDP sockets, tun disabled), lets it run, then Stop()s it on a timer
// and asserts it actually closes. If the reader goroutines parked in recvmmsg
// don't wake on Close(), Wait() blocks forever and this fails with a goroutine
// dump instead of relying on a process signal to unstick them.
func TestControlStopClosesOnTimer(t *testing.T) {
l := test.NewLogger()
dir := t.TempDir()
before := time.Now().Add(-time.Hour)
after := time.Now().Add(time.Hour)
ca, _, caKey, caPEM := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, before, after, nil, nil, nil)
networks := []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")}
_, _, keyPEM, certPEM := cert_test.NewTestCert(cert.Version2, cert.Curve_CURVE25519, ca, caKey, "close-on-timer", before, after, networks, nil, nil)
caPath := filepath.Join(dir, "ca.pem")
certPath := filepath.Join(dir, "cert.pem")
keyPath := filepath.Join(dir, "key.pem")
require.NoError(t, os.WriteFile(caPath, caPEM, 0o600))
require.NoError(t, os.WriteFile(certPath, certPEM, 0o600))
require.NoError(t, os.WriteFile(keyPath, keyPEM, 0o600))
// tun disabled so no device/root is needed; routines: 2 so we exercise the
// multi-socket (SO_REUSEPORT) teardown, which is where dnclient runs.
configBody := fmt.Sprintf(`
pki:
ca: %s
cert: %s
key: %s
listen:
host: 127.0.0.1
port: 0
tun:
disabled: true
firewall:
outbound:
- port: any
proto: any
host: any
inbound:
- port: any
proto: any
host: any
routines: 2
`, caPath, certPath, keyPath)
require.NoError(t, os.WriteFile(filepath.Join(dir, "config.yml"), []byte(configBody), 0o600))
c := config.NewC(l)
require.NoError(t, c.Load(dir))
ctrl, err := nebula.Main(c, false, "close-on-timer", l, nil)
require.NoError(t, err)
require.NoError(t, ctrl.Start())
// Run like a live nebula, then close on a timer, exactly as dnclient does.
<-time.NewTimer(5 * time.Second).C
stopped := make(chan struct{})
go func() {
ctrl.Stop() // closes the udp sockets (shutdown(2)) and the tun
ctrl.Wait() // blocks until every reader goroutine has returned
close(stopped)
}()
select {
case <-stopped:
t.Log("nebula closed cleanly on timer")
case <-time.After(10 * time.Second):
buf := make([]byte, 1<<20)
n := runtime.Stack(buf, true)
t.Fatalf("nebula did NOT close within 10s of Stop(): a blocking reader never woke\n%s", buf[:n])
}
}
+1
View File
@@ -25,6 +25,7 @@ func newTestLighthouse() *LightHouse {
lighthouses := []netip.Addr{} lighthouses := []netip.Addr{}
staticList := map[netip.Addr]struct{}{} staticList := map[netip.Addr]struct{}{}
lh.localAddrsFn = func(*LocalAllowList) []netip.Addr { return nil }
lh.lighthouses.Store(&lighthouses) lh.lighthouses.Store(&lighthouses)
lh.staticList.Store(&staticList) lh.staticList.Store(&staticList)
+65
View File
@@ -2,16 +2,26 @@ package nebula
import ( import (
"encoding/json" "encoding/json"
"log/slog"
"sync" "sync"
"sync/atomic" "sync/atomic"
"github.com/slackhq/nebula/cert" "github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/handshake" "github.com/slackhq/nebula/handshake"
"github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/noiseutil" "github.com/slackhq/nebula/noiseutil"
) )
const ReplayWindow = 8192 const ReplayWindow = 8192
// sessionEpoch hands out a receiver-local ordinal to every ConnectionState at creation. The RX
// staging sort (overlay/batch) orders packets by (epoch, message counter). A re-handshake never
// rekeys an existing tunnel; it brings up a new hostinfo and ConnectionState with a counter space
// starting near zero, while the old tunnel keeps decrypting until torn down. During that cutover
// one flush batch can hold packets from both tunnels, and the epoch keeps the old tunnel's
// packets sorted first.
var sessionEpoch atomic.Uint64
type ConnectionState struct { type ConnectionState struct {
eKey noiseutil.CipherState eKey noiseutil.CipherState
dKey noiseutil.CipherState dKey noiseutil.CipherState
@@ -20,7 +30,10 @@ type ConnectionState struct {
initiator bool initiator bool
messageCounter atomic.Uint64 messageCounter atomic.Uint64
window *Bits window *Bits
decryptLock sync.Mutex
writeLock sync.Mutex writeLock sync.Mutex
// epoch is this session's sessionEpoch ordinal. Immutable after creation.
epoch uint64
} }
// newConnectionStateFromResult builds a fully-populated ConnectionState from a // newConnectionStateFromResult builds a fully-populated ConnectionState from a
@@ -35,6 +48,7 @@ func newConnectionStateFromResult(r *handshake.Result) *ConnectionState {
eKey: noiseutil.NewCipherState(r.EKey, r.Cipher), eKey: noiseutil.NewCipherState(r.EKey, r.Cipher),
dKey: noiseutil.NewCipherState(r.DKey, r.Cipher), dKey: noiseutil.NewCipherState(r.DKey, r.Cipher),
window: NewBits(ReplayWindow), window: NewBits(ReplayWindow),
epoch: sessionEpoch.Add(1),
} }
ci.messageCounter.Add(r.MessageIndex) ci.messageCounter.Add(r.MessageIndex)
for i := uint64(1); i <= r.MessageIndex; i++ { for i := uint64(1); i <= r.MessageIndex; i++ {
@@ -54,3 +68,54 @@ func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
func (cs *ConnectionState) Curve() cert.Curve { func (cs *ConnectionState) Curve() cert.Curve {
return cs.myCert.Curve() return cs.myCert.Curve()
} }
func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) ([]byte, error) {
cs.decryptLock.Lock()
result := cs.window.Check(l, messageCounter)
cs.decryptLock.Unlock()
if !result {
return nil, ErrAlreadySeen
}
out, err := cs.dKey.DecryptDanger(packet[header.Len:header.Len], packet[:header.Len], packet[header.Len:], messageCounter, nb)
if err != nil {
return nil, err
}
cs.decryptLock.Lock()
result = cs.window.Update(l, messageCounter)
cs.decryptLock.Unlock()
if !result {
return nil, ErrAlreadySeen
}
return out, nil
}
func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) error {
cs.decryptLock.Lock()
result := cs.window.Check(l, messageCounter)
cs.decryptLock.Unlock()
if !result {
return ErrAlreadySeen
}
// The entire body is sent as AD, not encrypted.
// The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value.
// The packet is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's
// otherwise malformed (meaning, there is no trailing 16 byte AEAD value), then this will result in at worst a 0-length slice
// which will gracefully fail in the DecryptDanger call.
signedPayload := packet[:len(packet)-cs.dKey.Overhead()]
signatureValue := packet[len(packet)-cs.dKey.Overhead():]
_, err := cs.dKey.DecryptDanger(nil, signedPayload, signatureValue, messageCounter, nb)
if err != nil {
return err
}
cs.decryptLock.Lock()
result = cs.window.Update(l, messageCounter)
cs.decryptLock.Unlock()
if !result {
return ErrAlreadySeen
}
return nil
}
+9 -1
View File
@@ -53,6 +53,7 @@ type Control struct {
statsStart func() statsStart func()
dnsStart func() dnsStart func()
lighthouseStart func() lighthouseStart func()
networkChangeStart func(rebind func())
connectionManagerStart func(context.Context) connectionManagerStart func(context.Context)
} }
@@ -104,6 +105,9 @@ func (c *Control) Start() error {
if c.dnsStart != nil { if c.dnsStart != nil {
go c.dnsStart() go c.dnsStart()
} }
if c.networkChangeStart != nil {
go c.networkChangeStart(c.RebindUDPServer)
}
if c.connectionManagerStart != nil { if c.connectionManagerStart != nil {
go c.connectionManagerStart(c.ctx) go c.connectionManagerStart(c.ctx)
} }
@@ -198,7 +202,11 @@ func (c *Control) RebindUDPServer() {
return return
} }
_ = c.f.outside.Rebind() // A failure here means we are likely still pinned to the interface we came up on, so the rest of this is
// unlikely to help. Say so instead of silently carrying on as if we rebound.
if err := c.f.outside.Rebind(); err != nil {
c.l.Error("Failed to rebind udp socket", "error", err)
}
// Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0 // Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0
c.f.lightHouse.SendUpdate() c.f.lightHouse.SendUpdate()
+4 -4
View File
@@ -78,7 +78,7 @@ func newReadyControl(t *testing.T) (*Control, *fakeDevice, *fakeConn) {
inside: dev, inside: dev,
outside: conn, outside: conn,
writers: []udp.Conn{conn}, writers: []udp.Conn{conn},
batchers: make([]batch.RxBatcher, 1), batchers: make([]*batch.MultiCoalescer, 1),
routines: 1, routines: 1,
hostMap: newHostMap(l), hostMap: newHostMap(l),
lightHouse: lh, lightHouse: lh,
@@ -148,8 +148,8 @@ func (c *fakeConn) Rebind() error { c.rebinds++; ret
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil } func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
func (c *fakeConn) ListenOut(_ udp.EncReader, _ func()) error { return nil } func (c *fakeConn) ListenOut(_ udp.EncReader, _ func()) error { return nil }
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil } func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
func (c *fakeConn) WriteBatch(_ [][]byte, _ []netip.AddrPort, _ []byte) error { func (c *fakeConn) WriteBatch(bufs [][]byte, _ []netip.AddrPort) (int, error) {
return nil return len(bufs), nil
} }
func (c *fakeConn) ReloadConfig(_ *config.C) {} func (c *fakeConn) ReloadConfig(_ *config.C) {}
func (c *fakeConn) SupportsMultipleReaders() bool { return true } func (c *fakeConn) SupportsMultipleReaders() bool { return true }
@@ -177,7 +177,7 @@ func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
inside: dev, inside: dev,
outside: conn, outside: conn,
writers: []udp.Conn{conn}, writers: []udp.Conn{conn},
batchers: make([]batch.RxBatcher, 2), batchers: make([]*batch.MultiCoalescer, 2),
routines: 2, routines: 2,
l: test.NewLogger(), l: test.NewLogger(),
} }
+13 -1
View File
@@ -108,7 +108,19 @@ func (c *Control) GetVpnAddrs() []netip.Addr {
} }
func (c *Control) GetUDPAddr() netip.AddrPort { func (c *Control) GetUDPAddr() netip.AddrPort {
return c.f.outside.(*udp.TesterConn).Addr return c.f.outside.(*udp.TesterConn).GetAddr()
}
// SetUDPAddr moves this node to a new underlay address, standing in for a laptop waking up on a different
// network. Register the new address with the router as well or nothing will route back.
func (c *Control) SetUDPAddr(addr netip.AddrPort) {
c.f.outside.(*udp.TesterConn).SetAddr(addr)
}
// SetLocalAddrsFn replaces underlay address discovery so a test can advertise its simulated address instead of
// whatever this machine's NICs happen to be. Call it before Start, SendUpdate reads it from the update worker.
func (c *Control) SetLocalAddrsFn(fn func(*LocalAllowList) []netip.Addr) {
c.f.lightHouse.localAddrsFn = fn
} }
func (c *Control) KillPendingTunnel(vpnIp netip.Addr) bool { func (c *Control) KillPendingTunnel(vpnIp netip.Addr) bool {
+187
View File
@@ -0,0 +1,187 @@
// Package cpupick chooses which CPUs the tun reader threads pin to when the
// operator has not chosen for us (tun.cpu_affinity). The stock spread —
// allowed[i] for routine i — has two failure modes this package exists to fix:
//
// - every co-located nebula starts its spread at allowed[0], so N instances
// on one box stack their readers onto the same cores, and allowed[0] is
// usually CPU 0, the core housekeeping and default IRQ affinity already
// favor;
// - on heterogeneous CPUs (ARM big.LITTLE, Intel P/E hybrids, AMD compact
// cores) low IDs are not necessarily fast cores, and pinning an encrypt
// thread to an efficiency core caps that queue's throughput.
//
// Default instead returns a preference-ordered pin list: the allowed set
// filtered to performance cores (when the platform distinguishes them and
// enough remain for every routine), confined to a single NUMA node and spread
// across distinct physical cores when the topology permits, CPU 0's physical
// core demoted to last resort, and the order rotated by a stable per-instance
// key so co-located instances spread instead of stacking.
package cpupick
import (
"log/slog"
"github.com/slackhq/nebula/util"
)
// topology is the slice of machine layout arrange consults: the NUMA node
// and the physical core behind each candidate CPU, plus which core CPU 0
// lives on (zeroCore, -1 when unknown — tracked separately because CPU 0's
// SMT sibling deserves demotion even when CPU 0 itself isn't a candidate).
// Probed from sysfs on Linux; flatTopology stands in when the platform can't
// say, which turns every topology rule into a no-op rather than a wrong
// answer.
type topology struct {
nodeOf map[int]int
coreOf map[int]int
zeroCore int
}
// flatTopology places every CPU on node 0 and on a physical core of its own.
func flatTopology(cpus []int) topology {
t := topology{
nodeOf: make(map[int]int, len(cpus)),
coreOf: make(map[int]int, len(cpus)),
zeroCore: -1,
}
for i, c := range cpus {
t.nodeOf[c] = 0
t.coreOf[c] = i
if c == 0 {
t.zeroCore = i
}
}
return t
}
// Default computes the pin order for `routines` tun readers. key is any
// stable per-instance value; the bound UDP port is ideal — distinct across
// co-located instances, stable across restarts so benchmark runs stay
// comparable. Returns nil when there is nothing useful to say (no affinity
// support on this platform, lookup failure); callers keep their existing
// fallback spread.
func Default(routines int, key uint64, l *slog.Logger) []int {
allowed, err := util.AllowedCPUs()
if err != nil || len(allowed) == 0 {
return nil
}
perf, signal := perfCPUs(allowed)
cands := pickCandidates(allowed, perf, routines)
if len(cands) == 0 {
return nil
}
if len(perf) < routines {
signal = ""
}
cpus := arrange(cands, readTopology(cands), routines, splitmix64(key))
if l != nil {
l.Info("chose default pin CPUs for tun readers",
"cpus", cpus[:min(routines, len(cpus))],
"perfSignal", signal)
}
return cpus
}
// pickCandidates applies the enough-for-everyone guard: a perf filter that
// leaves fewer candidates than routines is discarded — giving every reader
// its own (possibly slow) core beats stacking two readers on a fast one.
func pickCandidates(allowed, perf []int, routines int) []int {
if len(perf) < routines {
return allowed
}
return perf
}
// arrange turns the candidate set into the final pin order:
//
// 1. NUMA: when at least one node holds enough candidates for every
// routine, confine to one such node, chosen by the instance hash. The
// readers share hostmap and cipher state, so splitting one instance
// across nodes taxes every packet — and co-located instances that hash
// to different nodes stop competing entirely. When no node is big
// enough, span nodes rather than stack readers.
// 2. Rotate the preferred candidates by the hash so instances spread.
// 3. SMT: emit one thread per physical core before any of their siblings —
// two encrypt threads on one core split its execution units. Siblings
// still follow for the routines > cores case.
// 4. CPU 0's whole physical core goes last: housekeeping and default IRQ
// noise on CPU 0 bleeds into its SMT sibling too. Within that tail the
// sibling precedes CPU 0 itself, which only catches the bleed-through.
//
// The rotation happens before the SMT pass so each instance's one-per-core
// walk also starts at a different core, and CPU 0's core is excluded from
// the rotation so no hash value can put it back at the front.
func arrange(cands []int, topo topology, routines int, h uint64) []int {
byNode := map[int][]int{}
var nodes []int
for _, c := range cands {
n := topo.nodeOf[c]
if _, ok := byNode[n]; !ok {
nodes = append(nodes, n)
}
byNode[n] = append(byNode[n], c)
}
var eligible []int
for _, n := range nodes {
if len(byNode[n]) >= routines {
eligible = append(eligible, n)
}
}
if len(eligible) > 0 {
cands = byNode[eligible[int(h%uint64(len(eligible)))]]
}
// Split off CPU 0's core: its siblings tail the list, CPU 0 tails them.
preferred := make([]int, 0, len(cands))
var zeroTail []int
hasZero := false
for _, c := range cands {
switch {
case c == 0:
hasZero = true
case topo.zeroCore >= 0 && topo.coreOf[c] == topo.zeroCore:
zeroTail = append(zeroTail, c)
default:
preferred = append(preferred, c)
}
}
if hasZero {
zeroTail = append(zeroTail, 0)
}
if len(preferred) == 0 {
return zeroTail // CPU 0's core is all we have
}
// The node pick consumed the low hash bits; rotate by the high ones so
// the two choices stay independent.
off := int((h >> 32) % uint64(len(preferred)))
rot := make([]int, 0, len(preferred))
rot = append(rot, preferred[off:]...)
rot = append(rot, preferred[:off]...)
seenCore := make(map[int]bool, len(rot))
out := make([]int, 0, len(cands))
var siblings []int
for _, c := range rot {
g := topo.coreOf[c]
if seenCore[g] {
siblings = append(siblings, c)
continue
}
seenCore[g] = true
out = append(out, c)
}
out = append(out, siblings...)
out = append(out, zeroTail...)
return out
}
// splitmix64 decorrelates instance keys before the selection modulos: ports
// on one box often share spacing (4242/4243, or round steps like +1000) that
// raw key%len arithmetic would fold onto the same offset.
func splitmix64(x uint64) uint64 {
x += 0x9e3779b97f4a7c15
x = (x ^ (x >> 30)) * 0xbf58476d1ce4e5b9
x = (x ^ (x >> 27)) * 0x94d049bb133111eb
return x ^ (x >> 31)
}
+171
View File
@@ -0,0 +1,171 @@
package cpupick
import (
"slices"
"testing"
)
// pairTopo builds a topology where consecutive candidate pairs are SMT
// siblings: (cpus[0],cpus[1]) share a core, (cpus[2],cpus[3]) the next, ...
// All CPUs land on node 0.
func pairTopo(cpus []int) topology {
t := topology{
nodeOf: make(map[int]int, len(cpus)),
coreOf: make(map[int]int, len(cpus)),
zeroCore: -1,
}
for i, c := range cpus {
t.nodeOf[c] = 0
t.coreOf[c] = i / 2
if c == 0 {
t.zeroCore = i / 2
}
}
return t
}
func TestArrangeDemotesZeroForEveryKey(t *testing.T) {
candidates := []int{0, 1, 2, 3, 4, 5, 6, 7}
for key := range uint64(64) {
got := arrange(candidates, flatTopology(candidates), 4, splitmix64(key))
if len(got) != len(candidates) {
t.Fatalf("key %d: len=%d want %d", key, len(got), len(candidates))
}
if got[0] == 0 {
t.Errorf("key %d: CPU 0 at the front: %v", key, got)
}
if got[len(got)-1] != 0 {
t.Errorf("key %d: CPU 0 not demoted to last: %v", key, got)
}
sorted := slices.Clone(got)
slices.Sort(sorted)
if !slices.Equal(sorted, candidates) {
t.Errorf("key %d: not a permutation: %v", key, got)
}
}
}
func TestArrangeDemotesZeroSiblings(t *testing.T) {
// Pairs (0,1),(2,3),(4,5),(6,7): CPU 0's core — 0 and its sibling 1 —
// must tail the list, sibling ahead of 0 itself.
candidates := []int{0, 1, 2, 3, 4, 5, 6, 7}
for key := range uint64(64) {
got := arrange(candidates, pairTopo(candidates), 2, splitmix64(key))
n := len(got)
if got[n-1] != 0 || got[n-2] != 1 {
t.Fatalf("key %d: tail = %v, want [... 1 0]", key, got)
}
}
}
func TestArrangeZeroSiblingWithoutZero(t *testing.T) {
// CPU 0 excluded (cpuset) but its sibling 1 remains: the sibling still
// tails the list when the topology knows which core CPU 0 lives on.
candidates := []int{1, 2, 3, 4, 5}
topo := pairTopo([]int{0, 1, 2, 3, 4, 5})
got := arrange(candidates, topo, 2, splitmix64(7))
if got[len(got)-1] != 1 {
t.Errorf("CPU 0's sibling not demoted: %v", got)
}
}
func TestArrangeRotatesByKey(t *testing.T) {
candidates := []int{1, 2, 3, 4, 5, 6, 7, 8}
seen := map[int]bool{}
for key := range uint64(64) {
seen[arrange(candidates, flatTopology(candidates), 4, splitmix64(key))[0]] = true
}
// 64 hashed keys over 8 slots must hit more than one starting CPU, or
// co-located instances would all stack again.
if len(seen) < 2 {
t.Errorf("rotation never varied across keys: %v", seen)
}
}
func TestArrangeStableForSameKey(t *testing.T) {
candidates := []int{0, 2, 4, 6}
topo := flatTopology(candidates)
a := arrange(candidates, topo, 2, splitmix64(4242))
b := arrange(candidates, topo, 2, splitmix64(4242))
if !slices.Equal(a, b) {
t.Errorf("same key ordered differently: %v vs %v", a, b)
}
}
func TestArrangeZeroOnly(t *testing.T) {
if got := arrange([]int{0}, flatTopology([]int{0}), 1, splitmix64(7)); !slices.Equal(got, []int{0}) {
t.Errorf("sole CPU 0 must survive: %v", got)
}
}
func TestArrangeSMTSiblingsLast(t *testing.T) {
// Pairs (1,2),(3,4),(5,6),(7,8): the first four picks must cover four
// distinct physical cores before any sibling repeats.
candidates := []int{1, 2, 3, 4, 5, 6, 7, 8}
topo := pairTopo(candidates)
for key := range uint64(16) {
got := arrange(candidates, topo, 4, splitmix64(key))
seen := map[int]bool{}
for _, c := range got[:4] {
g := topo.coreOf[c]
if seen[g] {
t.Fatalf("key %d: sibling before all cores covered: %v", key, got)
}
seen[g] = true
}
}
}
func TestArrangeNUMAConfinesToOneNode(t *testing.T) {
// Two nodes of four; both fit routines=3, so the result must sit
// entirely inside one of them, and the hash must pick both across keys.
candidates := []int{1, 2, 3, 4, 10, 11, 12, 13}
topo := flatTopology(candidates)
for _, c := range []int{10, 11, 12, 13} {
topo.nodeOf[c] = 1
}
nodesSeen := map[int]bool{}
for key := range uint64(32) {
got := arrange(candidates, topo, 3, splitmix64(key))
if len(got) != 4 {
t.Fatalf("key %d: not confined to one node: %v", key, got)
}
n := topo.nodeOf[got[0]]
for _, c := range got {
if topo.nodeOf[c] != n {
t.Fatalf("key %d: spans nodes: %v", key, got)
}
}
nodesSeen[n] = true
}
if len(nodesSeen) != 2 {
t.Errorf("hash never spread instances across nodes: %v", nodesSeen)
}
}
func TestArrangeNUMASpansWhenNoNodeFits(t *testing.T) {
candidates := []int{1, 2, 3, 4, 10, 11, 12, 13}
topo := flatTopology(candidates)
for _, c := range []int{10, 11, 12, 13} {
topo.nodeOf[c] = 1
}
got := arrange(candidates, topo, 6, splitmix64(1))
if len(got) != len(candidates) {
t.Errorf("undersized nodes must span, got %v", got)
}
}
func TestPickCandidates(t *testing.T) {
allowed := []int{0, 1, 2, 3, 4, 5, 6, 7}
perf := []int{4, 5}
// Enough perf cores for every routine: only they are used.
if got := pickCandidates(allowed, perf, 2); !slices.Equal(got, perf) {
t.Errorf("perf filter not applied: %v", got)
}
// Perf filter too small for the routine count: discarded, everyone
// gets their own core from the full allowed set.
if got := pickCandidates(allowed, perf, 4); !slices.Equal(got, allowed) {
t.Errorf("undersized perf filter not discarded: %v", got)
}
}
+154
View File
@@ -0,0 +1,154 @@
//go:build linux
package cpupick
import (
"fmt"
"os"
"path/filepath"
"strconv"
"strings"
)
// capacityKeepPct is the cpu_capacity admission threshold, relative to the
// fastest allowed core. LITTLE cores are normalized to ~250-400 of the big
// core's 1024 while mid cores sit at ~75%+, so half of max separates little
// from the rest without splitting prime from mid on three-tier parts.
const capacityKeepPct = 50
// freqKeepPct is the cpuinfo_max_freq admission threshold. Favored-core
// turbo skew is 2-4% and ARM mid-vs-prime ~12%, while E-cores, LITTLE
// cores, and AMD compact cores all sit >= 20% below their siblings' max.
const freqKeepPct = 85
// perfCPUs partitions allowed into the subset that are "performance" cores,
// consulting (in order of authority):
//
// 1. cpu_capacity — arch_topology's normalized per-CPU capacity, exposed on
// arm/arm64/riscv; the scheduler's own view of big vs LITTLE.
// 2. /sys/devices/cpu_core/cpus — the Intel hybrid P-core PMU mask, present
// only on P/E parts (x86 has no cpu_capacity) and naming P cores outright.
// 3. cpuinfo_max_freq — the cross-vendor fallback; catches AMD compact
// cores, which neither of the above covers.
//
// Returns allowed unchanged (signal "") when nothing distinguishes the
// cores: homogeneous parts, VMs without cpufreq, sysfs unavailable.
func perfCPUs(allowed []int) ([]int, string) {
return perfCPUsFrom("/sys/devices/system/cpu", "/sys/devices/cpu_core/cpus", allowed)
}
func perfCPUsFrom(cpuDir, intelCoreMask string, allowed []int) ([]int, string) {
if cpus, ok := byPerCPUValue(cpuDir, "cpu_capacity", allowed, capacityKeepPct); ok {
return cpus, "cpu_capacity"
}
if cpus, ok := byIntelCoreMask(intelCoreMask, allowed); ok {
return cpus, "intel_core_pmu"
}
if cpus, ok := byPerCPUValue(cpuDir, "cpufreq/cpuinfo_max_freq", allowed, freqKeepPct); ok {
return cpus, "max_freq"
}
return allowed, ""
}
// byPerCPUValue keeps the allowed CPUs whose per-CPU sysfs value is at least
// keepPct percent of the maximum across allowed. Inconclusive (ok=false)
// when any CPU is missing the file or when every value is equal.
func byPerCPUValue(cpuDir, file string, allowed []int, keepPct int) ([]int, bool) {
vals := make([]int, len(allowed))
minV, maxV := 0, 0
for i, cpu := range allowed {
v, err := readIntFile(filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), file))
if err != nil {
return nil, false
}
vals[i] = v
if i == 0 || v < minV {
minV = v
}
if v > maxV {
maxV = v
}
}
if minV == maxV {
return nil, false // homogeneous by this signal; try the next one
}
keep := make([]int, 0, len(allowed))
for i, cpu := range allowed {
if vals[i]*100 >= maxV*keepPct {
keep = append(keep, cpu)
}
}
return keep, true
}
// byIntelCoreMask keeps the allowed CPUs named by the hybrid P-core PMU
// mask. Inconclusive when the file is absent (non-hybrid x86, other arches)
// or no allowed CPU is in the mask (the process was deliberately confined
// to E-cores; nothing useful to prefer within that).
func byIntelCoreMask(maskPath string, allowed []int) ([]int, bool) {
b, err := os.ReadFile(maskPath)
if err != nil {
return nil, false
}
set, err := parseCPUList(strings.TrimSpace(string(b)))
if err != nil || len(set) == 0 {
return nil, false
}
pcore := make(map[int]bool, len(set))
for _, c := range set {
pcore[c] = true
}
keep := make([]int, 0, len(allowed))
for _, cpu := range allowed {
if pcore[cpu] {
keep = append(keep, cpu)
}
}
if len(keep) == 0 {
return nil, false
}
return keep, true
}
// parseCPUList decodes the kernel's cpulist format ("0-7,16-23", "3") into
// individual CPU IDs. Empty input yields an empty list.
func parseCPUList(s string) ([]int, error) {
if s == "" {
return nil, nil
}
var out []int
for part := range strings.SplitSeq(s, ",") {
part = strings.TrimSpace(part)
if part == "" {
continue
}
lo, hi, isRange := strings.Cut(part, "-")
a, err := strconv.Atoi(lo)
if err != nil {
return nil, fmt.Errorf("bad cpulist entry %q: %w", part, err)
}
if !isRange {
out = append(out, a)
continue
}
b, err := strconv.Atoi(hi)
if err != nil {
return nil, fmt.Errorf("bad cpulist entry %q: %w", part, err)
}
if b < a || b-a > 8192 {
return nil, fmt.Errorf("bad cpulist range %q", part)
}
for v := a; v <= b; v++ {
out = append(out, v)
}
}
return out, nil
}
func readIntFile(path string) (int, error) {
b, err := os.ReadFile(path)
if err != nil {
return 0, err
}
return strconv.Atoi(strings.TrimSpace(string(b)))
}
+163
View File
@@ -0,0 +1,163 @@
//go:build linux
package cpupick
import (
"fmt"
"os"
"path/filepath"
"slices"
"testing"
)
// fakeSysfs builds a cpuDir tree with the given per-CPU file values.
// A nil map for a file means "file absent on every CPU".
func fakeSysfs(t *testing.T, capacity, maxFreq map[int]int) string {
t.Helper()
dir := t.TempDir()
write := func(cpu int, rel string, v int) {
p := filepath.Join(dir, fmt.Sprintf("cpu%d", cpu), rel)
if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(p, fmt.Appendf(nil, "%d\n", v), 0o644); err != nil {
t.Fatal(err)
}
}
for cpu, v := range capacity {
write(cpu, "cpu_capacity", v)
}
for cpu, v := range maxFreq {
write(cpu, "cpufreq/cpuinfo_max_freq", v)
}
return dir
}
func writeCoreMask(t *testing.T, mask string) string {
t.Helper()
p := filepath.Join(t.TempDir(), "cpus")
if err := os.WriteFile(p, []byte(mask+"\n"), 0o644); err != nil {
t.Fatal(err)
}
return p
}
func TestPerfCPUsBigLittleCapacity(t *testing.T) {
// 4 big (1024) + 4 LITTLE (~290): capacity is authoritative on ARM.
dir := fakeSysfs(t, map[int]int{
0: 1024, 1: 1024, 2: 1024, 3: 1024,
4: 290, 5: 290, 6: 290, 7: 290,
}, nil)
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3, 4, 5, 6, 7})
if signal != "cpu_capacity" {
t.Fatalf("signal = %q", signal)
}
if !slices.Equal(got, []int{0, 1, 2, 3}) {
t.Errorf("got %v", got)
}
}
func TestPerfCPUsThreeTierKeepsMid(t *testing.T) {
// prime (1024) + mid (~780) + little (~280): 50% keeps prime+mid.
dir := fakeSysfs(t, map[int]int{
0: 280, 1: 280, 2: 280, 3: 280,
4: 780, 5: 780, 6: 780,
7: 1024,
}, nil)
got, _ := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3, 4, 5, 6, 7})
if !slices.Equal(got, []int{4, 5, 6, 7}) {
t.Errorf("got %v", got)
}
}
func TestPerfCPUsIntelHybridMask(t *testing.T) {
// No cpu_capacity on x86; the P-core PMU mask decides.
dir := fakeSysfs(t, nil, nil)
mask := writeCoreMask(t, "0-7")
got, signal := perfCPUsFrom(dir, mask, []int{0, 1, 2, 3, 8, 9, 10, 11})
if signal != "intel_core_pmu" {
t.Fatalf("signal = %q", signal)
}
if !slices.Equal(got, []int{0, 1, 2, 3}) {
t.Errorf("got %v", got)
}
}
func TestPerfCPUsIntelMaskDisjointFallsThrough(t *testing.T) {
// Confined to E-cores only: the mask can't help, and equal freqs below
// mean nothing else distinguishes them either -> allowed unchanged.
dir := fakeSysfs(t, nil, map[int]int{8: 4300000, 9: 4300000})
mask := writeCoreMask(t, "0-7")
got, signal := perfCPUsFrom(dir, mask, []int{8, 9})
if signal != "" || !slices.Equal(got, []int{8, 9}) {
t.Errorf("got %v signal %q", got, signal)
}
}
func TestPerfCPUsMaxFreqCompactCores(t *testing.T) {
// AMD-style compact cores: no capacity, no Intel mask; 3.3 vs 5.7 GHz.
dir := fakeSysfs(t, nil, map[int]int{
0: 5700000, 1: 5700000, 2: 3300000, 3: 3300000,
})
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3})
if signal != "max_freq" {
t.Fatalf("signal = %q", signal)
}
if !slices.Equal(got, []int{0, 1}) {
t.Errorf("got %v", got)
}
}
func TestPerfCPUsFavoredCoreSkewKept(t *testing.T) {
// Turbo Boost Max favored cores run a few percent hot; they must not
// shrink the candidate set to one or two cores.
dir := fakeSysfs(t, nil, map[int]int{
0: 5800000, 1: 5700000, 2: 5700000, 3: 5600000,
})
got, _ := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2, 3})
if !slices.Equal(got, []int{0, 1, 2, 3}) {
t.Errorf("favored-core skew filtered CPUs: %v", got)
}
}
func TestPerfCPUsHomogeneousInconclusive(t *testing.T) {
dir := fakeSysfs(t, nil, map[int]int{0: 3000000, 1: 3000000})
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1})
if signal != "" || !slices.Equal(got, []int{0, 1}) {
t.Errorf("got %v signal %q", got, signal)
}
}
func TestPerfCPUsNoSysfs(t *testing.T) {
dir := t.TempDir()
got, signal := perfCPUsFrom(dir, filepath.Join(dir, "nope"), []int{0, 1, 2})
if signal != "" || !slices.Equal(got, []int{0, 1, 2}) {
t.Errorf("got %v signal %q", got, signal)
}
}
func TestParseCPUList(t *testing.T) {
cases := []struct {
in string
want []int
wantErr bool
}{
{"0-3", []int{0, 1, 2, 3}, false},
{"0-1,16-17", []int{0, 1, 16, 17}, false},
{"5", []int{5}, false},
{"", nil, false},
{"3-1", nil, true},
{"a-b", nil, true},
{"1,x", nil, true},
}
for _, c := range cases {
got, err := parseCPUList(c.in)
if (err != nil) != c.wantErr {
t.Errorf("%q: err=%v wantErr=%v", c.in, err, c.wantErr)
continue
}
if !c.wantErr && !slices.Equal(got, c.want) {
t.Errorf("%q: got %v want %v", c.in, got, c.want)
}
}
}
+10
View File
@@ -0,0 +1,10 @@
//go:build !linux
package cpupick
// perfCPUs is Linux-only sysfs walking; elsewhere report "no distinction".
// Default already returns nil off-Linux (util.AllowedCPUs has no answer
// there), so this exists to keep the package compiling everywhere.
func perfCPUs(allowed []int) ([]int, string) {
return allowed, ""
}
+118
View File
@@ -0,0 +1,118 @@
//go:build linux
package cpupick
import (
"fmt"
"os"
"path/filepath"
"strconv"
"strings"
)
// readTopology probes the NUMA node and physical-core layout of cpus from
// sysfs. Anything sysfs won't say degrades toward flatTopology: an unknown
// node becomes node 0, an unknown core becomes a core of its own — either
// way the corresponding arrange rule becomes a no-op instead of a wrong
// answer.
func readTopology(cpus []int) topology {
return readTopologyFrom("/sys/devices/system/node", "/sys/devices/system/cpu", cpus)
}
func readTopologyFrom(nodeDir, cpuDir string, cpus []int) topology {
coreOf, zeroCore := coreGroups(cpuDir, cpus)
return topology{
nodeOf: numaNodes(nodeDir, cpus),
coreOf: coreOf,
zeroCore: zeroCore,
}
}
// numaNodes maps each cpu to its NUMA node via
// /sys/devices/system/node/nodeN/cpulist. CPUs no node claims (or no node
// dirs at all: VMs, non-NUMA kernels) land on node 0.
func numaNodes(nodeDir string, cpus []int) map[int]int {
out := make(map[int]int, len(cpus))
for _, c := range cpus {
out[c] = 0
}
entries, err := os.ReadDir(nodeDir)
if err != nil {
return out
}
want := make(map[int]bool, len(cpus))
for _, c := range cpus {
want[c] = true
}
for _, e := range entries {
id, ok := strings.CutPrefix(e.Name(), "node")
if !ok {
continue
}
n, err := strconv.Atoi(id)
if err != nil {
continue // has_cpu, possible, ... share the prefix
}
b, err := os.ReadFile(filepath.Join(nodeDir, e.Name(), "cpulist"))
if err != nil {
continue
}
list, err := parseCPUList(strings.TrimSpace(string(b)))
if err != nil {
continue
}
for _, c := range list {
if want[c] {
out[c] = n
}
}
}
return out
}
// coreGroups maps each cpu to a dense physical-core id derived from its
// (physical_package_id, core_id) pair — core_id alone repeats across
// sockets. CPUs whose topology files are unreadable get a core of their own.
// The second return is the group id of the core CPU 0 lives on, or -1 when
// that can't be determined; CPU 0's own files are consulted even when 0 is
// not a candidate, so its SMT siblings are recognized under cpusets that
// exclude CPU 0 itself.
func coreGroups(cpuDir string, cpus []int) (map[int]int, int) {
type pkgCore struct{ pkg, core int }
pairOf := func(cpu int) (pkgCore, bool) {
topoDir := filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), "topology")
pkg, err1 := readIntFile(filepath.Join(topoDir, "physical_package_id"))
core, err2 := readIntFile(filepath.Join(topoDir, "core_id"))
if err1 != nil || err2 != nil {
return pkgCore{}, false
}
return pkgCore{pkg, core}, true
}
ids := map[pkgCore]int{}
out := make(map[int]int, len(cpus))
next := 0
for _, cpu := range cpus {
k, ok := pairOf(cpu)
if !ok {
out[cpu] = next
next++
continue
}
id, ok := ids[k]
if !ok {
id = next
next++
ids[k] = id
}
out[cpu] = id
}
zeroCore := -1
if k, ok := pairOf(0); ok {
if id, ok := ids[k]; ok {
zeroCore = id
}
}
return out, zeroCore
}
+111
View File
@@ -0,0 +1,111 @@
//go:build linux
package cpupick
import (
"fmt"
"os"
"path/filepath"
"testing"
)
// fakeTopoSysfs builds nodeDir/cpuDir trees. nodes maps node id -> cpulist
// string; cores maps cpu -> (package, core) pair.
func fakeTopoSysfs(t *testing.T, nodes map[int]string, cores map[int][2]int) (string, string) {
t.Helper()
base := t.TempDir()
nodeDir := filepath.Join(base, "node")
cpuDir := filepath.Join(base, "cpu")
for n, list := range nodes {
d := filepath.Join(nodeDir, fmt.Sprintf("node%d", n))
if err := os.MkdirAll(d, 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(d, "cpulist"), []byte(list+"\n"), 0o644); err != nil {
t.Fatal(err)
}
}
for cpu, pc := range cores {
d := filepath.Join(cpuDir, fmt.Sprintf("cpu%d", cpu), "topology")
if err := os.MkdirAll(d, 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(d, "physical_package_id"), fmt.Appendf(nil, "%d\n", pc[0]), 0o644); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(d, "core_id"), fmt.Appendf(nil, "%d\n", pc[1]), 0o644); err != nil {
t.Fatal(err)
}
}
return nodeDir, cpuDir
}
func TestReadTopology(t *testing.T) {
// Two nodes; SMT pairs (0,4),(1,5) on node 0 and (2,6),(3,7) on node 1.
// core_id repeats across packages on purpose: the pair must disambiguate.
nodeDir, cpuDir := fakeTopoSysfs(t,
map[int]string{0: "0-1,4-5", 1: "2-3,6-7"},
map[int][2]int{
0: {0, 0}, 4: {0, 0}, 1: {0, 1}, 5: {0, 1},
2: {1, 0}, 6: {1, 0}, 3: {1, 1}, 7: {1, 1},
})
cpus := []int{0, 1, 2, 3, 4, 5, 6, 7}
topo := readTopologyFrom(nodeDir, cpuDir, cpus)
for _, c := range []int{0, 1, 4, 5} {
if topo.nodeOf[c] != 0 {
t.Errorf("cpu %d on node %d, want 0", c, topo.nodeOf[c])
}
}
for _, c := range []int{2, 3, 6, 7} {
if topo.nodeOf[c] != 1 {
t.Errorf("cpu %d on node %d, want 1", c, topo.nodeOf[c])
}
}
pairs := [][2]int{{0, 4}, {1, 5}, {2, 6}, {3, 7}}
for _, p := range pairs {
if topo.coreOf[p[0]] != topo.coreOf[p[1]] {
t.Errorf("siblings %v not grouped: %d vs %d", p, topo.coreOf[p[0]], topo.coreOf[p[1]])
}
}
if topo.coreOf[0] == topo.coreOf[2] {
t.Error("cross-package cores with equal core_id must not merge")
}
if topo.zeroCore != topo.coreOf[0] {
t.Errorf("zeroCore = %d, want %d", topo.zeroCore, topo.coreOf[0])
}
}
func TestReadTopologyZeroCoreWithoutZeroCandidate(t *testing.T) {
// CPU 0 is not a candidate (cpuset excludes it) but its sibling 4 is:
// zeroCore must still identify their shared core.
nodeDir, cpuDir := fakeTopoSysfs(t,
map[int]string{0: "0-7"},
map[int][2]int{0: {0, 0}, 4: {0, 0}, 1: {0, 1}, 5: {0, 1}})
topo := readTopologyFrom(nodeDir, cpuDir, []int{1, 4, 5})
if topo.zeroCore < 0 || topo.coreOf[4] != topo.zeroCore {
t.Errorf("zeroCore = %d, coreOf[4] = %d; sibling of CPU 0 not identified", topo.zeroCore, topo.coreOf[4])
}
if topo.coreOf[1] == topo.zeroCore {
t.Error("cpu 1 wrongly grouped with CPU 0's core")
}
}
func TestReadTopologyMissingSysfs(t *testing.T) {
base := t.TempDir()
cpus := []int{0, 1, 2}
topo := readTopologyFrom(filepath.Join(base, "nope"), filepath.Join(base, "also-nope"), cpus)
seen := map[int]bool{}
for _, c := range cpus {
if topo.nodeOf[c] != 0 {
t.Errorf("cpu %d node = %d, want 0", c, topo.nodeOf[c])
}
if seen[topo.coreOf[c]] {
t.Errorf("cpu %d shares a fallback core group", c)
}
seen[topo.coreOf[c]] = true
}
if topo.zeroCore != -1 {
t.Errorf("zeroCore = %d, want -1 when unknown", topo.zeroCore)
}
}
+10
View File
@@ -0,0 +1,10 @@
//go:build !linux
package cpupick
// readTopology has no sysfs to consult off Linux; the flat stand-in makes
// arrange's NUMA and SMT rules no-ops. Default is already nil off Linux
// (util.AllowedCPUs has no answer there) — this keeps the package compiling.
func readTopology(cpus []int) topology {
return flatTopology(cpus)
}
+16 -7
View File
@@ -97,8 +97,7 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
newAddr := getDnsServerAddr(c) newAddr := getDnsServerAddr(c)
d.serverMu.Lock() d.serverMu.Lock()
running := d.server running := d.server != nil
runningStarted := d.started
sameAddr := d.addr == newAddr sameAddr := d.addr == newAddr
d.addr = newAddr d.addr = newAddr
d.enabled.Store(enabled) d.enabled.Store(enabled)
@@ -112,7 +111,7 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
} }
if !enabled { if !enabled {
if running != nil { if running {
d.Stop() d.Stop()
} }
// Drop any records that accumulated while enabled; a later re-enable // Drop any records that accumulated while enabled; a later re-enable
@@ -121,12 +120,12 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
return nil return nil
} }
if running == nil { if !running {
// Was disabled (or never started); bring it up now. // Was disabled (or never started); bring it up now.
go d.Start() go d.Start()
} else if !sameAddr { } else if !sameAddr {
d.shutdownServer(running, runningStarted, "reload") // Stop clears the slot before shutting down, otherwise the Start below can find the dying server and refuse
// Old Start goroutine has now exited; bring up a fresh listener on the new address. d.Stop()
go d.Start() go d.Start()
} }
@@ -162,7 +161,9 @@ func (d *dnsServer) Start() {
started := make(chan struct{}) started := make(chan struct{})
d.serverMu.Lock() d.serverMu.Lock()
if d.ctx.Err() != nil { // Re-check enabled under the lock, a disable that raced our check above snapshots the slot under it too.
// Two reloads in quick succession can both spawn a Start, the loser would orphan the live listener past Stop
if d.ctx.Err() != nil || d.server != nil || !d.enabled.Load() {
d.serverMu.Unlock() d.serverMu.Unlock()
return return
} }
@@ -200,6 +201,14 @@ func (d *dnsServer) Start() {
close(started) close(started)
} }
// Release our slot, unless a reload already replaced us, so a dead listener can't block a future Start
d.serverMu.Lock()
if d.server == server {
d.server = nil
d.started = nil
}
d.serverMu.Unlock()
if err != nil { if err != nil {
d.l.Warn("Failed to run the DNS responder", "error", err) d.l.Warn("Failed to run the DNS responder", "error", err)
} }
+206 -4
View File
@@ -194,14 +194,51 @@ func TestDnsServer_reload_initial_serveDnsWithoutLighthouse(t *testing.T) {
} }
func TestDnsServer_reload_sameAddr_noOp(t *testing.T) { func TestDnsServer_reload_sameAddr_noOp(t *testing.T) {
port := freeUDPPort(t)
ds, c := newTestDnsServer(t) ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", "0", true, true) setDnsConfig(c, "127.0.0.1", port, true, true)
require.NoError(t, ds.reload(c, true)) require.NoError(t, ds.reload(c, true))
// No server running yet, no addr change. Reload should not spawn anything.
go ds.Start()
waitForBind(t, ds)
ds.serverMu.Lock()
before := ds.server
ds.serverMu.Unlock()
require.NotNil(t, before)
// Same address, so the running listener must be left alone rather than rebuilt under live queries
require.NoError(t, ds.reload(c, false)) require.NoError(t, ds.reload(c, false))
assert.True(t, ds.enabled.Load()) assert.True(t, ds.enabled.Load())
assert.Nil(t, ds.server)
ds.serverMu.Lock()
after := ds.server
ds.serverMu.Unlock()
assert.Same(t, before, after, "a same-address reload must not restart the listener")
ds.Stop()
}
// The branch the old sameAddr test was accidentally hitting: enabled with nothing running means reload starts it.
func TestDnsServer_reload_whenNotRunning_starts(t *testing.T) {
port := freeUDPPort(t)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", port, true, true)
// initial only records config, it never starts anything
require.NoError(t, ds.reload(c, true))
ds.serverMu.Lock()
assert.Nil(t, ds.server, "the initial reload must not start a listener")
ds.serverMu.Unlock()
require.NoError(t, ds.reload(c, false))
waitForBind(t, ds)
ds.serverMu.Lock()
assert.NotNil(t, ds.server, "a reload with nothing running should bring DNS up")
ds.serverMu.Unlock()
ds.Stop()
} }
func TestDnsServer_StartStop_lifecycle(t *testing.T) { func TestDnsServer_StartStop_lifecycle(t *testing.T) {
@@ -427,3 +464,168 @@ func waitFor(t *testing.T, cond func() bool) {
} }
t.Fatal("timed out waiting for condition") t.Fatal("timed out waiting for condition")
} }
// Two reloads in quick succession, or a HUP before Control.Start, can race two Starts at the same listener.
func TestDnsServer_Start_isIdempotent(t *testing.T) {
port := freeUDPPort(t)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", port, true, true)
require.NoError(t, ds.reload(c, true))
go ds.Start()
waitForBind(t, ds)
ds.serverMu.Lock()
first := ds.server
ds.serverMu.Unlock()
require.NotNil(t, first)
// If the second Start replaces the tracked server, Stop kills the wrong one and the port leaks
done := make(chan struct{})
go func() {
ds.Start()
close(done)
}()
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("second Start never returned")
}
ds.serverMu.Lock()
second := ds.server
ds.serverMu.Unlock()
assert.Same(t, first, second, "a second Start must not replace the running server")
// The real proof, after Stop the port must actually be free
ds.Stop()
waitFor(t, func() bool {
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
if err != nil {
return false
}
_ = pc.Close()
return true
})
}
// An address change must actually end up listening on the new port. Start's guard refuses when a server is already
// installed, so reload has to clear the slot before shutting the old one down.
func TestDnsServer_reload_addrChange_restarts(t *testing.T) {
first := freeUDPPort(t)
second := freeUDPPort(t)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", first, true, true)
require.NoError(t, ds.reload(c, true))
go ds.Start()
waitForBind(t, ds)
// Cycle a few times, the failure this guards against depends on which goroutine wins serverMu
for i := range 8 {
want := second
if i%2 == 1 {
want = first
}
setDnsConfig(c, "127.0.0.1", want, true, true)
require.NoError(t, ds.reload(c, false))
waitForBind(t, ds)
ds.serverMu.Lock()
srv := ds.server
ds.serverMu.Unlock()
require.NotNil(t, srv, "reload left DNS down instead of restarting it")
require.Equal(t, "127.0.0.1:"+want, srv.Addr, "reload should be serving the new address")
}
// Land back on second so the port assertions below are meaningful
setDnsConfig(c, "127.0.0.1", second, true, true)
require.NoError(t, ds.reload(c, false))
waitForBind(t, ds)
// The old port must be released and the new one actually held
waitFor(t, func() bool {
pc, err := net.ListenPacket("udp", "127.0.0.1:"+first)
if err != nil {
return false
}
_ = pc.Close()
return true
})
_, err := net.ListenPacket("udp", "127.0.0.1:"+second)
require.Error(t, err, "the new address should be bound by the DNS responder")
ds.Stop()
}
// A listener that dies on its own must release the slot, or a later same-addr reload sees it as running and no-ops.
func TestDnsServer_Start_bindFailure_releasesSlot(t *testing.T) {
port := freeUDPPort(t)
blocker, err := net.ListenPacket("udp", "127.0.0.1:"+port)
require.NoError(t, err)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", port, true, true)
require.NoError(t, ds.reload(c, true))
ds.Start() // returns once the bind fails
ds.serverMu.Lock()
assert.Nil(t, ds.server, "a listener that failed to bind must not stay parked in the slot")
ds.serverMu.Unlock()
// With the slot released, a reload can retry once the port frees up
require.NoError(t, blocker.Close())
require.NoError(t, ds.reload(c, false))
waitForBind(t, ds)
ds.serverMu.Lock()
assert.NotNil(t, ds.server, "a same-addr reload should retry after a failed bind")
ds.serverMu.Unlock()
ds.Stop()
}
// A disable that lands while Start is between its unlocked check and the guard must not leave a listener behind.
func TestDnsServer_Start_refusesWhenDisabledUnderLock(t *testing.T) {
port := freeUDPPort(t)
ds, c := newTestDnsServer(t)
setDnsConfig(c, "127.0.0.1", port, true, true)
require.NoError(t, ds.reload(c, true))
require.True(t, ds.enabled.Load())
// Holding serverMu parks Start on the lock, the only way to land the disable in that window on purpose
ds.serverMu.Lock()
done := make(chan struct{})
go func() {
ds.Start()
close(done)
}()
select {
case <-done:
ds.serverMu.Unlock()
t.Fatal("Start returned early, the test never exercised the window")
case <-time.After(time.Millisecond * 100):
}
// The disable reload's critical section. It sees nothing running, so it never calls Stop.
ds.enabled.Store(false)
ds.serverMu.Unlock()
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("Start never returned")
}
ds.serverMu.Lock()
assert.Nil(t, ds.server, "Start must not install a listener a disable already cancelled")
ds.serverMu.Unlock()
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
require.NoError(t, err, "an orphaned listener is still holding the port")
_ = pc.Close()
}
+64
View File
@@ -725,6 +725,70 @@ func TestReestablishRelays(t *testing.T) {
} }
func TestRelayHandshakeOverDisestablishedEntry(t *testing.T) {
t.Parallel()
// If them tears down the tunnel while me keeps Established relay state, me's next
// handshake flows through the relay with no fresh CreateRelayRequest and lands on
// them's Disestablished terminal relay entry. them must re-establish that entry, or
// its first transmit deletes its only relay and the tunnel is born transmit-dead:
// them can receive but every send is silently dropped.
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them ", "10.128.0.2/24", m{"relay": m{"use_relays": true}})
// Teach my how to get to the relay and that their can be reached via the relay
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
// Build a router so we don't have to reason who gets which packet
r := router.NewR(t, myControl, relayControl, theirControl)
defer r.RenderFlow()
// Start the servers
myControl.Start()
relayControl.Start()
theirControl.Start()
t.Log("Trigger a handshake from me to them via the relay")
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
p := r.RouteForAllUntilTxTun(theirControl)
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
oldIdx := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false).LocalIndex
t.Log("Close the tunnel on them only, marking their relay entry Disestablished")
theirControl.CloseTunnel(myVpnIpNet[0].Addr(), true)
t.Log("Re-handshake from me, riding the still-Established relay state")
myControl.ReHandshake(theirVpnIpNet[0].Addr())
for {
h := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false)
if h != nil && h.LocalIndex != oldIdx && h.RemoteIndex != 0 {
break
}
r.RouteForAllExitFunc(func(*udp.Packet, *nebula.Control) router.ExitType {
return router.RouteAndExit
})
}
hAtThem := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
require.NotNil(t, hAtThem, "them should have completed the relayed handshake")
require.Equal(t, []netip.Addr{relayVpnIpNet[0].Addr()}, hAtThem.CurrentRelaysToMe, "them should know a relay for the new tunnel")
t.Log("Send from them to me; their only relay entry must survive the transmit")
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
require.Never(t, func() bool {
h := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
return h == nil || len(h.CurrentRelaysToMe) == 0
}, time.Second, 10*time.Millisecond, "them deleted its only relay entry; the tunnel is permanently transmit-dead")
p = r.RouteForAllUntilTxTun(myControl)
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
r.RenderHostmaps("Final hostmaps", myControl, relayControl, theirControl)
}
func TestStage1RaceRelays(t *testing.T) { func TestStage1RaceRelays(t *testing.T) {
t.Parallel() t.Parallel()
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay //NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
+225
View File
@@ -0,0 +1,225 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
import (
"net/netip"
"testing"
"time"
"github.com/slackhq/nebula"
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/cert_test"
"github.com/slackhq/nebula/e2e/router"
"github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/udp"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// reportedAddrs is what the lighthouse would hand a peer asking where vpnAddr is.
func reportedAddrs(t *testing.T, lh *nebula.Control, vpnAddr netip.Addr) []netip.AddrPort {
t.Helper()
cm := lh.QueryLighthouse(vpnAddr)
if cm == nil {
return nil
}
var out []netip.AddrPort
for _, c := range *cm {
out = append(out, c.Reported...)
out = append(out, c.Learned...)
}
return out
}
// waitForLighthouseMsg routes until a lighthouse message lands on lh, or gives up. Reports whether one arrived.
func waitForLighthouseMsg(t *testing.T, r *router.R, lh *nebula.Control, wait time.Duration) bool {
t.Helper()
h := &header.H{}
return r.RouteForAllExitFuncOrTimeout(wait, func(p *udp.Packet, c *nebula.Control) router.ExitType {
if c != lh {
return router.KeepRouting
}
// Punches are a single byte and never parse, they are just not what we are after
if err := h.Parse(p.Data); err != nil {
return router.KeepRouting
}
if h.Type == header.LightHouse {
return router.RouteAndExit
}
return router.KeepRouting
})
}
// A laptop that changes networks has to tell the lighthouse promptly, otherwise the lighthouse keeps handing peers
// the old address and their punches land nowhere. On a long lighthouse interval the only thing that closes that
// window is the rebind, which on darwin the network change monitor drives. The e2e build compiles the monitor out,
// so we call RebindUDPServer directly, which is the same thing the monitor does.
func TestRebindSendsLighthouseUpdate(t *testing.T) {
t.Parallel()
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
"lighthouse": m{"am_lighthouse": true},
})
// 600s interval, so nothing scheduled can send an update during this test. A rebind is the only thing that can.
myControl, _, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
"lighthouse": m{
"hosts": []any{lhVpnIpNet[0].Addr().String()},
"interval": 600,
},
"static_host_map": m{
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
},
})
r := router.NewR(t, lhControl, myControl)
defer r.RenderFlow()
lhControl.Start()
myControl.Start()
// Let the startup registration finish, then clear everything it left behind
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
r.RouteFor(time.Millisecond * 400)
// Nothing should be talking to the lighthouse on its own now
require.False(t, waitForLighthouseMsg(t, r, lhControl, time.Millisecond*200),
"nothing should reach the lighthouse before the rebind")
myControl.RebindUDPServer()
assert.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5),
"a rebind should push an update to the lighthouse rather than waiting out the interval")
lhControl.Stop()
myControl.Stop()
}
// The other half of a rebind: every live tunnel requeries the lighthouse on its next send. That query is what makes
// the lighthouse tell the peer to punch toward our new address, which is the part that actually revives a tunnel
// whose remote NAT state died while we were on a different network.
func TestRebindRequeriesPeersOnNextSend(t *testing.T) {
t.Parallel()
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
"lighthouse": m{"am_lighthouse": true},
})
lhCfg := m{
"lighthouse": m{
"hosts": []any{lhVpnIpNet[0].Addr().String()},
"interval": 600,
// Without this the peers advertise this machine's real addresses and then try to punch at them,
// which the router has no route for.
"local_allow_list": m{
"10.0.0.0/24": true,
"::/0": false,
},
},
"static_host_map": m{
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
},
}
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", lhCfg)
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.3/24", lhCfg)
r := router.NewR(t, lhControl, myControl, theirControl)
defer r.RenderFlow()
lhControl.Start()
myControl.Start()
theirControl.Start()
r.RouteFor(time.Millisecond * 500)
// Point the peers at each other directly, this test is about the rebind and not about lighthouse discovery
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("initial")))
r.RouteFor(time.Second)
require.NotNil(t, myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false), "expected a tunnel to them")
r.RouteFor(time.Millisecond * 300)
// Assert on what the peer sees rather than on lighthouse traffic. A query for them makes the lighthouse send
// them a punch notification, which is the whole point. Our own update to the lighthouse sends them nothing,
// so this cannot be satisfied by the update the rebind itself pushes.
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("quiet")))
require.False(t, waitForLighthouseMsg(t, r, theirControl, time.Millisecond*300),
"an ordinary send should not requery the lighthouse")
myControl.RebindUDPServer()
r.RouteFor(time.Millisecond * 300) // let the update the rebind itself sends pass by
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("after rebind")))
assert.True(t, waitForLighthouseMsg(t, r, theirControl, time.Second*5),
"the first send after a rebind should requery the lighthouse, which then tells the peer to punch at us")
lhControl.Stop()
myControl.Stop()
theirControl.Stop()
}
// The scenario this whole thing exists for: a laptop sleeps at the office and wakes up at home on a new address.
// Until it tells the lighthouse, the lighthouse keeps handing peers the office address, so their punches land
// nowhere and the tunnel stays dead. On a long interval the rebind is the only thing that closes that window.
func TestRebindAdvertisesNewAddressAfterMove(t *testing.T) {
t.Parallel()
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
"lighthouse": m{"am_lighthouse": true},
})
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
"lighthouse": m{
"hosts": []any{lhVpnIpNet[0].Addr().String()},
"interval": 600,
},
"static_host_map": m{
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
},
})
// Advertise wherever we currently are rather than this machine's real NICs, read fresh each time so a move
// is picked up.
myControl.SetLocalAddrsFn(func(*nebula.LocalAllowList) []netip.Addr {
return []netip.Addr{myControl.GetUDPAddr().Addr()}
})
r := router.NewR(t, lhControl, myControl)
defer r.RenderFlow()
lhControl.Start()
myControl.Start()
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
r.RouteFor(time.Millisecond * 400)
require.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), myUdpAddr,
"the lighthouse should know the address we started on")
// Wake up somewhere else
newAddr := netip.MustParseAddrPort("10.0.0.99:4242")
myControl.SetUDPAddr(newAddr)
r.AddRoute(newAddr.Addr(), newAddr.Port(), myControl)
// Nothing has told the lighthouse, and with interval 600 nothing scheduled will
r.RouteFor(time.Millisecond * 400)
require.NotContains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
"the lighthouse should still be handing out the old address before the rebind")
myControl.RebindUDPServer()
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an update after the rebind")
r.RouteFor(time.Millisecond * 400)
assert.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
"after the rebind the lighthouse should hand peers our new address")
lhControl.Stop()
myControl.Stop()
}
+136
View File
@@ -0,0 +1,136 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
import (
"testing"
"time"
"github.com/slackhq/nebula"
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/cert_test"
"github.com/slackhq/nebula/e2e/router"
"github.com/slackhq/nebula/udp"
)
// TestRecoveryTiming measures how long a tunnel takes to come back after the peer stops accepting our traffic,
// which is what a laptop waking on a new network looks like from the peer's side: its NAT has no state for where
// we are now, so everything we send disappears.
//
// It is a measurement, not a pass/fail assertion. Recovery is timed to the moment the peer punches back at us,
// since that is when its NAT opens and the tunnel is usable again.
//
// go test -tags e2e_testing -v -run TestRecoveryTiming ./e2e/
func TestRecoveryTiming(t *testing.T) {
for _, tc := range []struct {
name string
rebind bool
}{
{"no trigger", false},
{"rebind counter", true},
} {
t.Run(tc.name, func(t *testing.T) {
d, lost := measureRecovery(t, tc.rebind)
t.Logf("RESULT %-16s recovered in %-9v (%d packets lost)", tc.name, d.Round(time.Millisecond), lost)
})
}
}
// measureRecovery returns how long until the peer punched back, and how many of our packets died meanwhile. When
// rebind is true we call RebindUDPServer once the tunnel goes dark, which is what the darwin network change
// monitor does and what iOS has always done. When false, nothing tells nebula anything is wrong.
func measureRecovery(t *testing.T, rebind bool) (time.Duration, int) {
t.Helper()
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
"lighthouse": m{"am_lighthouse": true},
})
peerCfg := m{
"lighthouse": m{
"hosts": []any{lhVpnIpNet[0].Addr().String()},
"interval": 600,
"local_allow_list": m{
"10.0.0.0/24": true,
"::/0": false,
},
},
"static_host_map": m{
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
},
}
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", peerCfg)
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.3/24", peerCfg)
r := router.NewR(t, lhControl, myControl, theirControl)
defer r.RenderFlow()
defer func() {
lhControl.Stop()
myControl.Stop()
theirControl.Stop()
}()
lhControl.Start()
myControl.Start()
theirControl.Start()
r.RouteFor(time.Millisecond * 500)
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("establish")))
r.RouteFor(time.Second)
if myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false) == nil {
t.Fatal("failed to establish the tunnel we are measuring")
}
r.RouteFor(time.Millisecond * 500)
// From here the peer's NAT has no state for us, everything we send it disappears
start := time.Now()
blackholed := 0
var recovered time.Duration
if rebind {
myControl.RebindUDPServer()
}
// Keep the tun busy the way someone retrying a stalled connection would
stop := make(chan struct{})
defer close(stop)
go func() {
tick := time.NewTicker(time.Millisecond * 200)
defer tick.Stop()
for {
select {
case <-stop:
return
case <-tick.C:
myControl.InjectTunPacket(BuildTunUDPPacket(
theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("retry")))
}
}
}()
r.RouteForAllExitFuncOrTimeout(time.Second*30, func(p *udp.Packet, c *nebula.Control) router.ExitType {
if c == theirControl && p.From == myControl.GetUDPAddr() {
blackholed++
return router.Drop
}
// The peer reaching us directly is the moment its NAT opened, whether that is a punch or a handshake
if c == myControl && p.From == theirUdpAddr {
recovered = time.Since(start)
return router.RouteAndExit
}
return router.KeepRouting
})
if recovered == 0 {
t.Fatalf("no recovery within 30s (%d packets blackholed)", blackholed)
}
return recovered, blackholed
}
+140 -21
View File
@@ -114,6 +114,28 @@ type packet struct {
packet *udp.Packet packet *udp.Packet
tun bool // a packet pulled off a tun device tun bool // a packet pulled off a tun device
rx bool // the packet was received by a udp device rx bool // the packet was received by a udp device
// h is the nebula header, parsed once when the packet is recorded. parseErr says why there isn't one, which
// the flow log reports rather than hiding. Punchy sends a single byte, so an unparseable packet is normal.
h header.H
parseErr error
}
// fromAddr and toAddr are the addresses this packet actually travelled between. Reading them off the control
// instead would misreport the whole history once a test moves a node. Tun packets are synthesized without
// addresses, so they fall back to the control.
func (p *packet) fromAddr() netip.AddrPort {
if p.tun || !p.packet.From.IsValid() {
return p.from.GetUDPAddr()
}
return p.packet.From
}
func (p *packet) toAddr() netip.AddrPort {
if p.tun || !p.packet.To.IsValid() {
return p.to.GetUDPAddr()
}
return p.packet.To
} }
func (p *packet) WasReceived() { func (p *packet) WasReceived() {
@@ -131,6 +153,9 @@ const (
ExitNow ExitType = 1 ExitNow ExitType = 1
// RouteAndExit routes this packet and exits immediately afterwards // RouteAndExit routes this packet and exits immediately afterwards
RouteAndExit ExitType = 2 RouteAndExit ExitType = 2
// Drop discards this packet without delivering it and keeps routing. Use it to simulate a blackhole, such as
// a restrictive NAT refusing traffic from an address it has not seen.
Drop ExitType = 3
) )
type ExitFunc func(packet *udp.Packet, receiver *nebula.Control) ExitType type ExitFunc func(packet *udp.Packet, receiver *nebula.Control) ExitType
@@ -141,7 +166,9 @@ type ExitFunc func(packet *udp.Packet, receiver *nebula.Control) ExitType
func NewR(t testing.TB, controls ...*nebula.Control) *R { func NewR(t testing.TB, controls ...*nebula.Control) *R {
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
if err := os.MkdirAll("mermaid", 0755); err != nil { // t.Name() contains a slash for subtests, so the flow log can land in a nested directory
fn := filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name()))
if err := os.MkdirAll(filepath.Dir(fn), 0755); err != nil {
panic(err) panic(err)
} }
@@ -152,7 +179,7 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
outNat: make(map[outNatKey]netip.AddrPort), outNat: make(map[outNatKey]netip.AddrPort),
flow: []flowEntry{}, flow: []flowEntry{},
ignoreFlows: []ignoreFlow{}, ignoreFlows: []ignoreFlow{},
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())), fn: fn,
t: t, t: t,
cancelRender: cancel, cancelRender: cancel,
} }
@@ -249,7 +276,7 @@ func (r *R) renderFlow() {
continue continue
} }
addr := e.packet.from.GetUDPAddr() addr := e.packet.fromAddr()
if _, ok := participants[addr]; ok { if _, ok := participants[addr]; ok {
continue continue
} }
@@ -268,7 +295,6 @@ func (r *R) renderFlow() {
} }
// Print packets // Print packets
h := &header.H{}
for _, e := range r.flow { for _, e := range r.flow {
if e.packet == nil { if e.packet == nil {
//fmt.Fprintf(f, " note over %s: %s\n", strings.Join(participantsVals, ", "), e.note) //fmt.Fprintf(f, " note over %s: %s\n", strings.Join(participantsVals, ", "), e.note)
@@ -280,21 +306,22 @@ func (r *R) renderFlow() {
fmt.Fprintln(f, r.formatUdpPacket(p)) fmt.Fprintln(f, r.formatUdpPacket(p))
} else { } else {
if err := h.Parse(p.packet.Data); err != nil {
panic(err)
}
line := "--x" line := "--x"
if p.rx { if p.rx {
line = "->>" line = "->>"
} }
fmt.Fprintf(f, detail := fmt.Sprintf("%s(%s), index %v, counter: %v",
" %s%s%s: %s(%s), index %v, counter: %v\n", p.h.TypeName(), p.h.SubTypeName(), p.h.RemoteIndex, p.h.MessageCounter)
normalizeName(p.from.GetUDPAddr().String()), if p.parseErr != nil {
detail = fmt.Sprintf("unparsed, %v (%d bytes)", p.parseErr, len(p.packet.Data))
}
fmt.Fprintf(f, " %s%s%s: %s\n",
normalizeName(p.fromAddr().String()),
line, line,
normalizeName(p.to.GetUDPAddr().String()), normalizeName(p.toAddr().String()),
h.TypeName(), h.SubTypeName(), h.RemoteIndex, h.MessageCounter, detail,
) )
} }
} }
@@ -408,21 +435,24 @@ func (r *R) unlockedInjectFlow(from, to *nebula.Control, p *udp.Packet, tun bool
r.renderHostmaps(fmt.Sprintf("Packet %v", len(r.flow))) r.renderHostmaps(fmt.Sprintf("Packet %v", len(r.flow)))
if len(r.ignoreFlows) > 0 {
var h header.H var h header.H
err := h.Parse(p.Data) var parseErr error
if err != nil { if !tun {
panic(err) parseErr = h.Parse(p.Data)
} }
// Decide before copying, the copy comes from a freelist and an ignored packet would never be released
for _, i := range r.ignoreFlows { for _, i := range r.ignoreFlows {
if !tun { if tun {
if i.messageType == h.Type && i.subType == h.Subtype { if i.tun.HasValue && i.tun.IsTrue {
return nil return nil
} }
} else if i.tun.HasValue && i.tun.IsTrue { continue
return nil
} }
// A packet we could not parse has no type to match against, so no rule can ignore it
if parseErr == nil && i.messageType == h.Type && i.subType == h.Subtype {
return nil
} }
} }
@@ -431,6 +461,8 @@ func (r *R) unlockedInjectFlow(from, to *nebula.Control, p *udp.Packet, tun bool
to: to, to: to,
packet: p.Copy(), packet: p.Copy(),
tun: tun, tun: tun,
h: h,
parseErr: parseErr,
} }
r.flow = append(r.flow, flowEntry{packet: fp}) r.flow = append(r.flow, flowEntry{packet: fp})
@@ -660,6 +692,10 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
p.Release() p.Release()
return return
case Drop:
// Record it so the flow log shows the attempt, but never hand it to the receiver
r.unlockedInjectFlow(sender, receiver, p, false)
case KeepRouting: case KeepRouting:
fp := r.unlockedInjectFlow(sender, receiver, p, false) fp := r.unlockedInjectFlow(sender, receiver, p, false)
receiver.InjectUDPPacket(p) receiver.InjectUDPPacket(p)
@@ -690,6 +726,85 @@ func (r *R) RouteUntilAfterMsgType(sender *nebula.Control, msgType header.Messag
}) })
} }
// RouteFor routes everything that shows up for the given duration and then returns. Use it to let a test settle
// deterministically rather than sleeping and hoping: a single FlushAll races a completing handshake, which queues
// more packets right behind it.
func (r *R) RouteFor(d time.Duration) {
r.RouteForAllExitFuncOrTimeout(d, func(*udp.Packet, *nebula.Control) ExitType {
return KeepRouting
})
}
// RouteForAllExitFuncOrTimeout is RouteForAllExitFunc with a deadline, reporting whether whatDo asked to exit
// before time ran out. The unbounded version blocks forever on a quiet network, so this is what a test needs to
// assert that something does NOT happen, or to route for a fixed settling period.
func (r *R) RouteForAllExitFuncOrTimeout(timeout time.Duration, whatDo ExitFunc) bool {
sc := make([]reflect.SelectCase, 0, len(r.controls)+1)
cm := make([]*nebula.Control, 0, len(r.controls))
for _, c := range r.controls {
sc = append(sc, reflect.SelectCase{
Dir: reflect.SelectRecv,
Chan: reflect.ValueOf(c.GetUDPTxChan()),
Send: reflect.Value{},
})
cm = append(cm, c)
}
timer := time.NewTimer(timeout)
defer timer.Stop()
sc = append(sc, reflect.SelectCase{
Dir: reflect.SelectRecv,
Chan: reflect.ValueOf(timer.C),
Send: reflect.Value{},
})
for {
x, rx, _ := reflect.Select(sc)
if x == len(cm) {
return false
}
r.Lock()
p := rx.Interface().(*udp.Packet)
receiver := r.getControl(cm[x].GetUDPAddr(), p.To, p)
if receiver == nil {
r.Unlock()
panic("Can't RouteForAllExitFuncOrTimeout for host: " + p.To.String())
}
e := whatDo(p, receiver)
switch e {
case ExitNow:
r.Unlock()
p.Release()
return true
case RouteAndExit:
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
receiver.InjectUDPPacket(p)
fp.WasReceived()
r.Unlock()
p.Release()
return true
case Drop:
// Record it so the flow log shows the attempt, but never hand it to the receiver
r.unlockedInjectFlow(cm[x], receiver, p, false)
case KeepRouting:
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
receiver.InjectUDPPacket(p)
fp.WasReceived()
default:
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
}
r.Unlock()
p.Release()
}
}
func (r *R) RouteForAllUntilAfterMsgTypeTo(receiver *nebula.Control, msgType header.MessageType, subType header.MessageSubType) { func (r *R) RouteForAllUntilAfterMsgTypeTo(receiver *nebula.Control, msgType header.MessageType, subType header.MessageSubType) {
h := &header.H{} h := &header.H{}
r.RouteForAllExitFunc(func(p *udp.Packet, r *nebula.Control) ExitType { r.RouteForAllExitFunc(func(p *udp.Packet, r *nebula.Control) ExitType {
@@ -782,6 +897,10 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
p.Release() p.Release()
return return
case Drop:
// Record it so the flow log shows the attempt, but never hand it to the receiver
r.unlockedInjectFlow(cm[x], receiver, p, false)
case KeepRouting: case KeepRouting:
fp := r.unlockedInjectFlow(cm[x], receiver, p, false) fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
receiver.InjectUDPPacket(p) receiver.InjectUDPPacket(p)
-188
View File
@@ -1,188 +0,0 @@
package nebula
import (
"encoding/binary"
"log/slog"
"testing"
"golang.org/x/net/ipv4"
)
func TestInnerECN(t *testing.T) {
cases := []struct {
name string
pkt []byte
want byte
}{
{"empty", nil, 0},
{"v4_NotECT", v4WithToS(0x00), 0x00},
{"v4_ECT0", v4WithToS(0x02), 0x02},
{"v4_ECT1", v4WithToS(0x01), 0x01},
{"v4_CE", v4WithToS(0x03), 0x03},
{"v4_DSCP_then_NotECT", v4WithToS(0x88 | 0x00), 0x00},
{"v4_DSCP_then_CE", v4WithToS(0x88 | 0x03), 0x03},
{"v6_NotECT", v6WithTC(0x00), 0x00},
{"v6_ECT0", v6WithTC(0x02), 0x02},
{"v6_CE", v6WithTC(0x03), 0x03},
{"v6_DSCP_then_CE", v6WithTC(0x88 | 0x03), 0x03},
{"unknown_version", []byte{0xa5, 0xff}, 0},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
got := innerECN(c.pkt)
if got != c.want {
t.Errorf("innerECN=0x%02x want 0x%02x", got, c.want)
}
})
}
}
// v4WithToS returns a 2-byte slice tall enough for innerECN: byte 0 carries
// version=4 in the high nibble, byte 1 is the full ToS so we exercise both
// the DSCP and ECN portions through the byte 1 mask.
func v4WithToS(tos byte) []byte {
return []byte{0x45, tos}
}
// v6WithTC builds a 2-byte slice that places a known traffic class value
// across bytes 0 (high nibble of TC) and 1 (low nibble of TC). innerECN
// extracts ECN as (b[1]>>4)&0x03, which corresponds to TC[1:0].
func v6WithTC(tc byte) []byte {
return []byte{0x60 | (tc>>4)&0x0f, (tc & 0x0f) << 4}
}
func TestApplyOuterECN(t *testing.T) {
silent := slog.New(slog.DiscardHandler)
hi := &HostInfo{}
// Build a v4 packet helper with a given inner ECN field.
v4 := func(innerECN byte) []byte {
// 20-byte minimal IPv4 header with ToS = innerECN (DSCP zeroed).
return []byte{
0x45, innerECN, 0, 28,
0, 0, 0x40, 0,
64, 6, 0, 0,
10, 0, 0, 1,
10, 0, 0, 2,
}
}
// Build a v6 packet helper with a given inner ECN field. ECN occupies
// TC[1:0] which sit at byte 1 mask 0x30.
v6 := func(innerECN byte) []byte {
// 40-byte minimal IPv6 header with TC[1:0] = innerECN.
pkt := make([]byte, 40)
pkt[0] = 0x60 // version=6, TC[7:4]=0
pkt[1] = (innerECN & 0x03) << 4 // TC[3:0]: low 2 bits = ECN, top 2 = DSCP-low (0)
return pkt
}
type cell struct {
outer byte
inner byte
wantECN byte
wantSame bool // expect inner unchanged (true => verify the byte didn't move)
}
// RFC 6040 normal-mode combine table. Only outer==CE causes mutation.
table := []cell{
{ecnNotECT, ecnNotECT, ecnNotECT, true},
{ecnNotECT, ecnECT0, ecnECT0, true},
{ecnNotECT, ecnECT1, ecnECT1, true},
{ecnNotECT, ecnCE, ecnCE, true},
{ecnECT0, ecnNotECT, ecnNotECT, true},
{ecnECT0, ecnECT0, ecnECT0, true},
{ecnECT0, ecnECT1, ecnECT1, true},
{ecnECT0, ecnCE, ecnCE, true},
{ecnECT1, ecnNotECT, ecnNotECT, true},
{ecnECT1, ecnECT0, ecnECT0, true},
{ecnECT1, ecnECT1, ecnECT1, true},
{ecnECT1, ecnCE, ecnCE, true},
{ecnCE, ecnNotECT, ecnNotECT, true}, // legacy: log, leave alone
{ecnCE, ecnECT0, ecnCE, false}, // CE folded in
{ecnCE, ecnECT1, ecnCE, false},
{ecnCE, ecnCE, ecnCE, true},
}
for _, c := range table {
t.Run("v4", func(t *testing.T) {
pkt := v4(c.inner)
applyOuterECN(pkt, c.outer, hi, silent)
got := pkt[1] & 0x03
if got != c.wantECN {
t.Errorf("v4 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
}
})
t.Run("v6", func(t *testing.T) {
pkt := v6(c.inner)
applyOuterECN(pkt, c.outer, hi, silent)
got := (pkt[1] >> 4) & 0x03
if got != c.wantECN {
t.Errorf("v6 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
}
})
}
}
// TestApplyOuterECN_IPv4ChecksumStaysValid guards against H1: folding an outer
// CE mark into the inner IPv4 ToS byte must keep the IPv4 header checksum valid.
// The passthrough emit paths write the packet verbatim, so a stale checksum
// turns an underlay congestion mark into packet loss at the receiver.
func TestApplyOuterECN_IPv4ChecksumStaysValid(t *testing.T) {
silent := slog.New(slog.DiscardHandler)
hi := &HostInfo{}
// 20-byte IPv4 header with DSCP=0x88 and inner ECN = ECT(0). Folding CE
// flips only the low two bits of the ToS byte while leaving DSCP intact.
pkt := []byte{
0x45, 0x88 | ecnECT0, 0, 40,
0x1c, 0x46, 0x40, 0x00,
64, 6, 0, 0,
10, 0, 0, 1,
10, 0, 0, 2,
}
// Stamp a correct header checksum before the fold.
binary.BigEndian.PutUint16(pkt[10:12], ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]))
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
t.Fatal("test setup: initial header checksum invalid")
}
applyOuterECN(pkt, ecnCE, hi, silent)
// CE folded in, DSCP preserved.
if got, want := pkt[1], byte(0x88|ecnCE); got != want {
t.Fatalf("ToS after fold = 0x%02x, want 0x%02x", got, want)
}
// The incremental RFC 1624 update must leave the checksum valid and equal
// to a full recompute over the mutated header.
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
t.Fatalf("IPv4 header checksum invalid after CE fold: 0x%04x", binary.BigEndian.Uint16(pkt[10:12]))
}
if got, want := binary.BigEndian.Uint16(pkt[10:12]), ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]); got != want {
t.Fatalf("checksum = 0x%04x, full recompute = 0x%04x", got, want)
}
}
// ipv4HeaderChecksum computes the RFC 1071 IPv4 header checksum over hdr,
// treating the checksum field (bytes 10:12) as zero.
func ipv4HeaderChecksum(hdr []byte) uint16 {
var sum uint32
for i := 0; i+1 < len(hdr); i += 2 {
if i == 10 {
continue // checksum field
}
sum += uint32(hdr[i])<<8 | uint32(hdr[i+1])
}
for sum > 0xffff {
sum = (sum >> 16) + (sum & 0xffff)
}
return ^uint16(sum)
}
// ipv4HeaderChecksumValid reports whether the stored checksum matches a fresh
// computation over the header.
func ipv4HeaderChecksumValid(hdr []byte) bool {
return binary.BigEndian.Uint16(hdr[10:12]) == ipv4HeaderChecksum(hdr)
}
+18 -20
View File
@@ -131,6 +131,9 @@ listen:
port: 4242 port: 4242
# Sets the max number of packets to pull from the kernel for each syscall (under systems that support recvmmsg) # Sets the max number of packets to pull from the kernel for each syscall (under systems that support recvmmsg)
# default is 64, does not support reload # default is 64, does not support reload
# Note: on Linux with UDP GRO (kernel 5.10+), each receive slot is sized for a full 64KiB coalesced
# superpacket, so the receive scratch is batch * 64KiB per listening socket (~4MiB per routine at the
# default of 64). Lower this to trade peak per-syscall throughput for memory on constrained hosts.
#batch: 64 #batch: 64
# Configure socket buffers for the udp side (outside), leave unset to use the system defaults. Values will be doubled by the kernel # Configure socket buffers for the udp side (outside), leave unset to use the system defaults. Values will be doubled by the kernel
# Default is net.core.rmem_default and net.core.wmem_default (/proc/sys/net/core/rmem_default and /proc/sys/net/core/rmem_default) # Default is net.core.rmem_default and net.core.wmem_default (/proc/sys/net/core/rmem_default and /proc/sys/net/core/rmem_default)
@@ -146,6 +149,14 @@ listen:
# Default true; set to false to leave WDF in charge of inbound decisions on the listener port. Not reloadable. # Default true; set to false to leave WDF in charge of inbound decisions on the listener port. Not reloadable.
#windows_bypass_wdf: true #windows_bypass_wdf: true
# On macOS only
# macOS scopes the udp socket to the interface it was created on, so moving between networks (wifi to wired,
# office to home) leaves Nebula sending out an interface that no longer has a route. When true, Nebula watches
# the routing socket and rebinds the listener once the change settles.
# iOS does not use this, the host app drives the same rebind itself.
# Default true. Not reloadable.
#rebind_on_network_change: true
# By default, Nebula replies to packets it has no tunnel for with a "recv_error" packet. This packet helps speed up reconnection # By default, Nebula replies to packets it has no tunnel for with a "recv_error" packet. This packet helps speed up reconnection
# in the case that Nebula on either side did not shut down cleanly. This response can be abused as a way to discover if Nebula is running # in the case that Nebula on either side did not shut down cleanly. This response can be abused as a way to discover if Nebula is running
# on a host though. This option lets you configure if you want to send "recv_error" packets always, never, or only to private network remotes. # on a host though. This option lets you configure if you want to send "recv_error" packets always, never, or only to private network remotes.
@@ -256,22 +267,19 @@ tun:
# Linux only. pin_threads pins each tun reader/encrypt OS thread to a single CPU. This keeps every goroutine's # Linux only. pin_threads pins each tun reader/encrypt OS thread to a single CPU. This keeps every goroutine's
# batched sends flowing through one XPS-selected NIC TX ring, so packets within a flow stay ordered on the wire # batched sends flowing through one XPS-selected NIC TX ring, so packets within a flow stay ordered on the wire
# instead of being sprayed across multiple TX rings and reordered. Not reloadable. # instead of being sprayed across multiple TX rings and reordered. Not reloadable. Coerced to false if routines <= 1.
#
# When cpu_affinity is unset, nebula picks CPUs that do NOT service any physical NIC's interrupts (read from
# /sys/class/net/*/device/msi_irqs and /proc/irq/*/effective_affinity_list): an encrypt thread pinned onto a core
# that also runs NAPI for a NIC RX queue fights the softirq for the core and collapses throughput for flows hashed
# to that queue. If the NIC's vectors blanket every allowed CPU (many drivers default to one queue per core) the
# avoidance logs and falls back to the old spread; narrow the NIC's queue/IRQ spread (e.g. `ethtool -X <dev>
# equal N`) or set cpu_affinity explicitly to benefit.
#pin_threads: true #pin_threads: true
# Linux only. cpu_affinity overrides which CPUs the tun reader threads pin to: a list of CPU IDs, one per routine # Linux only. cpu_affinity overrides which CPUs the tun reader threads pin to: a list of CPU IDs, one per routine
# (see the top-level `routines` setting). Lists shorter than `routines` are modulo-cycled across the queues; extra # (see the top-level `routines` setting). Lists shorter than `routines` are modulo-cycled across the queues; extra
# entries are ignored. IDs must be within the process's allowed CPU set, so this respects taskset / cgroup cpusets; # entries are ignored. IDs must be within the process's allowed CPU set, so this respects taskset / cgroup cpusets;
# a non-integer or not-allowed entry disables the override and falls back to spreading queues across the allowed # a non-integer or not-allowed entry disables the override and falls back to spreading queues across the allowed
# CPUs. Setting this disables the automatic NIC-IRQ avoidance described under pin_threads — prefer CPUs that don't # CPUs. Only meaningful while pin_threads is true. Not reloadable.
# service your underlay NIC's RX queue IRQs. Only meaningful while pin_threads is true. Not reloadable. # When unset, the default spread prefers performance cores on heterogeneous CPUs (ARM big.LITTLE, Intel P/E
# hybrids, AMD compact cores), keeps all readers on one NUMA node and on distinct physical cores when the
# topology allows (SMT siblings last), leaves CPU 0's physical core as a last resort, and rotates its starting
# point per instance (keyed by the bound UDP port) so co-located nebulas don't stack their readers onto the
# same cores.
#cpu_affinity: #cpu_affinity:
# - 2 # - 2
# - 4 # - 4
@@ -412,16 +420,6 @@ logging:
# This setting is reloadable # This setting is reloadable
#inactivity_timeout: 10m #inactivity_timeout: 10m
# ecn (default true) propagates ECN (Explicit Congestion Notification) across the tunnel per RFC 6040: the inner
# packet's ECN codepoint is copied onto the outer carrier header on encapsulation, and an outer CE ("congestion
# experienced") mark is folded back into the inner header on decapsulation. On linux it additionally stamps
# RTAX_FEATURE_ECN on the routes nebula installs, so the kernel actively negotiates ECN for connections to mesh
# prefixes. Disable this only when an underlay middlebox mangles or clears ECN bits unpredictably.
# This setting is reloadable, BUT flipping it at runtime only updates the datapath (the inner<->outer copy/combine).
# The RTAX_FEATURE_ECN flag on already-installed routes is NOT revisited on reload, so nebula must be restarted for
# the route half of this setting to take effect.
#ecn: true
# Nebula security group configuration # Nebula security group configuration
firewall: firewall:
# Action to take when a packet is not allowed by the firewall rules. # Action to take when a packet is not allowed by the firewall rules.
+9
View File
@@ -8,6 +8,15 @@ Before=sshd.service
Type=notify Type=notify
NotifyAccess=main NotifyAccess=main
SyslogIdentifier=nebula SyslogIdentifier=nebula
# Uncomment to run as an unprivileged user with only CAP_NET_ADMIN. Requires a
# nebula user that owns the config directory. Add CAP_NET_BIND_SERVICE to both
# lines if any listener (lighthouse DNS, listen.port, stats, sshd) binds <1024.
#User=nebula
#Group=nebula
#CapabilityBoundingSet=CAP_NET_ADMIN
#AmbientCapabilities=CAP_NET_ADMIN
ExecReload=/bin/kill -HUP $MAINPID ExecReload=/bin/kill -HUP $MAINPID
ExecStart=/usr/local/bin/nebula -config /etc/nebula/config.yml ExecStart=/usr/local/bin/nebula -config /etc/nebula/config.yml
Restart=always Restart=always
+9
View File
@@ -65,3 +65,12 @@ func (fp Packet) MarshalJSON() ([]byte, error) {
"Fragment": fp.Fragment, "Fragment": fp.Fragment,
}) })
} }
// ParsedPacket is a Packet plus the parse byproducts the RX path reuses
type ParsedPacket struct {
Packet
IPHdrLen int
// FragAny reports any fragmentation at all: MF flag or nonzero offset for IPv4, a fragment extension header for IPv6.
// Distinct from Packet.Fragment, which is true only for NON-FIRST fragments
FragAny bool
}
+6 -6
View File
@@ -1,6 +1,6 @@
module github.com/slackhq/nebula module github.com/slackhq/nebula
go 1.25.0 go 1.26.0
require ( require (
dario.cat/mergo v1.0.2 dario.cat/mergo v1.0.2
@@ -24,12 +24,12 @@ require (
github.com/vishvananda/netlink v1.3.1 github.com/vishvananda/netlink v1.3.1
go.uber.org/goleak v1.3.0 go.uber.org/goleak v1.3.0
go.yaml.in/yaml/v3 v3.0.4 go.yaml.in/yaml/v3 v3.0.4
golang.org/x/crypto v0.53.0 golang.org/x/crypto v0.54.0
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
golang.org/x/net v0.56.0 golang.org/x/net v0.57.0
golang.org/x/sync v0.21.0 golang.org/x/sync v0.22.0
golang.org/x/sys v0.46.0 golang.org/x/sys v0.47.0
golang.org/x/term v0.44.0 golang.org/x/term v0.45.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b
golang.zx2c4.com/wireguard/windows v1.0.1 golang.zx2c4.com/wireguard/windows v1.0.1
+10 -10
View File
@@ -162,8 +162,8 @@ golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACk
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4= golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto= golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw=
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio= golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk=
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY= golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc= golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY= golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
@@ -182,8 +182,8 @@ golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLL
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA= golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU= golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o= golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE=
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec= golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU=
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw= golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
@@ -191,8 +191,8 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM= golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
@@ -208,11 +208,11 @@ golang.org/x/sys v0.0.0-20210124154548-22da62e12c0c/go.mod h1:h1NjWce9XRLGQEsW7w
golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc= golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
golang.org/x/term v0.44.0/go.mod h1:7ze4MdzUzLXpSAoFP1H0bOI9aXDqveSvatT5vKcFh2Y= golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk= golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
+14 -4
View File
@@ -295,7 +295,13 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
hm.messageMetrics.Tx(header.Handshake, hh.machine.Subtype(), 1) hm.messageMetrics.Tx(header.Handshake, hh.machine.Subtype(), 1)
err := hm.outside.WriteTo(stage0, addr) err := hm.outside.WriteTo(stage0, addr)
if err != nil { if err != nil {
hostinfo.logger(hm.l).Error("Failed to send handshake message", // These repeat every attempt, so match the success log below and only shout when the remotes changed
level := slog.LevelDebug
if remotesHaveChanged {
level = slog.LevelError
}
hostinfo.logger(hm.l).Log(context.Background(), level, "Failed to send handshake message",
"udpAddr", addr, "udpAddr", addr,
"initiatorIndex", hostinfo.localIndexId, "initiatorIndex", hostinfo.localIndexId,
"handshake", hsFields, "handshake", hsFields,
@@ -529,8 +535,10 @@ func (hm *HandshakeManager) DeleteHostInfo(hostinfo *HostInfo) {
func (hm *HandshakeManager) unlockedDeleteHostInfo(hostinfo *HostInfo) { func (hm *HandshakeManager) unlockedDeleteHostInfo(hostinfo *HostInfo) {
for _, addr := range hostinfo.vpnAddrs { for _, addr := range hostinfo.vpnAddrs {
if cur, ok := hm.vpnIps[addr]; ok && cur.hostinfo == hostinfo {
delete(hm.vpnIps, addr) delete(hm.vpnIps, addr)
} }
}
if len(hm.vpnIps) == 0 { if len(hm.vpnIps) == 0 {
hm.vpnIps = map[netip.Addr]*HandshakeHostInfo{} hm.vpnIps = map[netip.Addr]*HandshakeHostInfo{}
@@ -967,7 +975,9 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
nb := make([]byte, 12, 12) nb := make([]byte, 12, 12)
out := make([]byte, mtu) out := make([]byte, mtu)
for _, cp := range hh.packetStore { for _, cp := range hh.packetStore {
//todo use a sendbatcher // TODO: use a SendBatch here. Each callback lands in
// sendNoMetrics -> WriteTo: one syscall per cached packet,
// where one sendmmsg could flush the whole store.
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out) cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
} }
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore))) f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
@@ -1078,8 +1088,8 @@ func (hm *HandshakeManager) sendHandshakeResponse(via ViaSender, msg []byte, hos
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0]) hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
// We received a valid handshake on this relay, so make sure the relay // We received a valid handshake on this relay, so make sure the relay
// state reflects that, in case it had been marked Disestablished. // state reflects that, in case it had been marked Disestablished.
via.relayHI.relayState.UpdateRelayForByIdxState(via.remoteIdx, Established) via.relayHI.relayState.UpdateRelayForByIdxState(via.relay.LocalIndex, Established)
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false) f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false, 0)
f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...) f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...)
} }
} }
+1 -1
View File
@@ -84,7 +84,7 @@ func (mw *mockEncWriter) SendMessageToVpnAddr(_ header.MessageType, _ header.Mes
return return
} }
func (mw *mockEncWriter) SendVia(_ *HostInfo, _ *Relay, _, _, _ []byte, _ bool) { func (mw *mockEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
return return
} }
+11 -6
View File
@@ -190,13 +190,18 @@ func SubTypeName(t MessageType, s MessageSubType) string {
} }
func IsValidSubType(t MessageType, s MessageSubType) bool { func IsValidSubType(t MessageType, s MessageSubType) bool {
if n, ok := subTypeMap[t]; ok { switch t {
if _, ok := (*n)[s]; ok { case Message:
return true return s == MessageNone || s == MessageRelay
} case Handshake:
} return s == HandshakeIXPSK0
case Test:
return s == TestReply || s == TestRequest
case Control, CloseTunnel, RecvError, LightHouse:
return s == 0
default:
return false return false
}
} }
// NewHeader turns bytes into a header // NewHeader turns bytes into a header
+51
View File
@@ -102,6 +102,57 @@ func TestTypeMap(t *testing.T) {
}, subTypeMap) }, subTypeMap)
} }
// mapIsValidSubType is the pre-refactor, map-driven definition of a valid
// subtype. IsValidSubType was reimplemented as an explicit switch; this keeps
// the original behavior around so we can prove the switch is equivalent to it.
func mapIsValidSubType(t MessageType, s MessageSubType) bool {
if n, ok := subTypeMap[t]; ok {
if _, ok := (*n)[s]; ok {
return true
}
}
return false
}
func TestIsValidSubType(t *testing.T) {
// Explicit intent table: documents exactly which subtypes are valid so the
// test stays meaningful even if both the switch and subTypeMap change.
assert.True(t, IsValidSubType(Message, MessageNone))
assert.True(t, IsValidSubType(Message, MessageRelay))
assert.False(t, IsValidSubType(Message, 2))
assert.True(t, IsValidSubType(Handshake, HandshakeIXPSK0))
// HandshakeXXPSK0 is defined but not a wire-valid subtype.
assert.False(t, IsValidSubType(Handshake, HandshakeXXPSK0))
assert.True(t, IsValidSubType(Test, TestRequest))
assert.True(t, IsValidSubType(Test, TestReply))
assert.False(t, IsValidSubType(Test, 2))
// These types only ever carry subtype 0.
for _, mt := range []MessageType{Control, CloseTunnel, RecvError, LightHouse} {
assert.True(t, IsValidSubType(mt, 0), "type %d subtype 0 should be valid", mt)
assert.False(t, IsValidSubType(mt, 1), "type %d subtype 1 should be invalid", mt)
}
// Unknown/unassigned types are never valid.
assert.False(t, IsValidSubType(99, 0))
// Exhaustive proof of equivalence with the original map-driven logic across
// the entire (type, subtype) input space.
for ti := 0; ti <= 0xff; ti++ {
for si := 0; si <= 0xff; si++ {
mt, mst := MessageType(ti), MessageSubType(si)
assert.Equalf(t, mapIsValidSubType(mt, mst), IsValidSubType(mt, mst),
"IsValidSubType(%d, %d) diverged from map-driven definition", ti, si)
}
}
// H method must delegate to the package function.
assert.True(t, (&H{Type: Test, Subtype: TestReply}).IsValidSubType())
assert.False(t, (&H{Type: Handshake, Subtype: HandshakeXXPSK0}).IsValidSubType())
}
func TestHeader_String(t *testing.T) { func TestHeader_String(t *testing.T) {
assert.Equal( assert.Equal(
t, t,
+11 -1
View File
@@ -287,7 +287,6 @@ type HostInfo struct {
type ViaSender struct { type ViaSender struct {
UdpAddr netip.AddrPort UdpAddr netip.AddrPort
relayHI *HostInfo // relayHI is the host info object of the relay relayHI *HostInfo // relayHI is the host info object of the relay
remoteIdx uint32 // remoteIdx is the index included in the header of the received packet
relay *Relay // relay contains the rest of the relay information, including the PeerIP of the host trying to communicate with us. relay *Relay // relay contains the rest of the relay information, including the PeerIP of the host trying to communicate with us.
IsRelayed bool // IsRelayed is true if the packet was sent through a relay IsRelayed bool // IsRelayed is true if the packet was sent through a relay
} }
@@ -544,6 +543,17 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
return final return final
} }
func (hm *HostMap) QueryIndexCached(index uint32, cache map[uint32]*HostInfo) *HostInfo {
if out, ok := cache[index]; ok {
return out
}
out := hm.QueryIndex(index)
if out != nil {
cache[index] = out
}
return out
}
func (hm *HostMap) QueryIndex(index uint32) *HostInfo { func (hm *HostMap) QueryIndex(index uint32) *HostInfo {
hm.RLock() hm.RLock()
if h, ok := hm.Indexes[index]; ok { if h, ok := hm.Indexes[index]; ok {
+31 -59
View File
@@ -15,7 +15,7 @@ import (
"github.com/slackhq/nebula/routing" "github.com/slackhq/nebula/routing"
) )
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packet, nb []byte, sendBatch batch.TxBatcher, 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 // borrowed: pkt.Bytes is owned by the originating tio.Queue and is
// only valid until the next Read on that queue. Every consumer below // only valid until the next Read on that queue. Every consumer below
// (parse, self-forward, handshake cache, sendInsideMessage) reads it // (parse, self-forward, handshake cache, sendInsideMessage) reads it
@@ -74,7 +74,7 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
return 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 // borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt
// bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket), // bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket),
// so retaining segments past the loop is safe. // so retaining segments past the loop is safe.
@@ -105,9 +105,9 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
return 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 { if dropReason == nil {
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch, rejectBuf, q) f.sendInsideMessage(hostinfo, pkt, nb, sendBatch)
} else { } else {
f.rejectInside(packet, rejectBuf, q) f.rejectInside(packet, rejectBuf, q)
if f.l.Enabled(context.Background(), slog.LevelDebug) { if f.l.Enabled(context.Background(), slog.LevelDebug) {
@@ -126,7 +126,6 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
c := ci.messageCounter.Add(1) c := ci.messageCounter.Add(1)
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c) out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
f.connectionManager.Out(hostinfo)
out, encErr := ci.eKey.EncryptDanger(out, out, seg, c, nb) out, encErr := ci.eKey.EncryptDanger(out, out, seg, c, nb)
if noiseutil.EncryptLockNeeded { if noiseutil.EncryptLockNeeded {
@@ -138,8 +137,7 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
"udpAddr", hostinfo.GetRemote(), "udpAddr", hostinfo.GetRemote(),
"counter", c, "counter", c,
) )
// Skip this segment; the rest of the superpacket can still // Skip this segment; the rest of the superpacket can still go out. TCP will retransmit anything we drop here.
// go out — TCP will retransmit anything we drop here.
return nil return nil
} }
@@ -151,16 +149,19 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
// later sendmmsg flush. Segmentation is fused with encryption here so the // later sendmmsg flush. Segmentation is fused with encryption here so the
// kernel-supplied superpacket bytes never get written into a separate // kernel-supplied superpacket bytes never get written into a separate
// scratch arena: SegmentSuperpacket builds each segment's plaintext in // scratch arena: SegmentSuperpacket builds each segment's plaintext in
// segScratch[:segLen] in turn, and we encrypt directly into a fresh // segScratch[:segLen] in turn, and we encrypt directly into a fresh SendBatch slot.
// SendBatch slot. func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []byte, sendBatch *batch.SendBatch) {
func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int) {
ci := hostinfo.ConnectionState ci := hostinfo.ConnectionState
if ci.eKey == nil { if ci.eKey == nil {
return return
} }
// One traffic-out mark covers every segment of the superpacket; doing it
// per segment in sendInsideEncrypt paid an atomic store up to ~45 extra
// times per TSO packet, inside writeLock under boring crypto.
f.connectionManager.Out(hostinfo)
remote := hostinfo.GetRemote() remote := hostinfo.GetRemote()
ecnEnabled := f.ecnEnabled.Load()
if hostinfo.lastRebindCount != f.rebindCount { if hostinfo.lastRebindCount != f.rebindCount {
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is //NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help. // finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
@@ -211,11 +212,7 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
return nil return nil
} }
var ecn byte sendBatch.Commit(toSend, relayHostInfo.GetRemote())
if ecnEnabled {
ecn = innerECN(seg)
}
sendBatch.Commit(toSend, relayHostInfo.GetRemote(), ecn)
return nil return nil
}) })
if err != nil { if err != nil {
@@ -233,36 +230,14 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
return nil return nil
} }
var ecn byte sendBatch.Commit(out, remote)
if ecnEnabled {
ecn = innerECN(seg)
}
sendBatch.Commit(out, remote, ecn)
return nil return nil
}) })
if err != nil { if err != nil {
hostinfo.logger(f.l).Error("Failed to segment superpacket for send", hostinfo.logger(f.l).Error("Failed to segment superpacket for send", "error", err)
"error", err,
)
} }
} }
// innerECN returns the 2-bit IP-level ECN codepoint of an inner IPv4 or IPv6
// packet, or 0 if pkt is too short or its IP version is unrecognized. Used at
// encap to copy the inner codepoint onto the outer carrier per RFC 6040.
func innerECN(pkt []byte) byte {
if len(pkt) < 2 {
return 0
}
switch pkt[0] >> 4 {
case 4:
return pkt[1] & 0x03
case 6:
return (pkt[1] >> 4) & 0x03
}
return 0
}
func (f *Interface) rejectInside(packet []byte, out []byte, q int) { func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
if !f.firewall.OutboundSendReject { if !f.firewall.OutboundSendReject {
return return
@@ -279,27 +254,30 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
} }
} }
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) { func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, rejectBuf []byte, q int) {
if !f.firewall.InboundSendReject { if !f.firewall.InboundSendReject {
return return
} }
out = iputil.CreateRejectPacket(packet, out) // split rejectBuf to make sure we have room to write the plaintext rejection, then encrypt it, without trampling anything
// we can't re-use packet, if we need to send an icmp reject, it won't be long enough.
half := len(rejectBuf) / 2
encryptBuf := rejectBuf[0:0:half] //the first half of rejectBuf's capacity, len set to 0
buildBuf := rejectBuf[half:]
out := iputil.CreateRejectPacket(packet, buildBuf)
if len(out) == 0 { if len(out) == 0 {
return return
} }
if len(out) > iputil.MaxRejectPacketSize { if len(out) > iputil.MaxRejectPacketSize {
if f.l.Enabled(context.Background(), slog.LevelInfo) { if f.l.Enabled(context.Background(), slog.LevelInfo) {
f.l.Info("rejectOutside: packet too big, not sending", f.l.Info("rejectOutside: packet too big, not sending", "packet", packet, "outPacket", out)
"packet", packet,
"outPacket", out,
)
} }
return return
} }
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, packet, q) f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, encryptBuf, q)
} }
// Handshake will attempt to initiate a tunnel with the provided vpn address. This is a no-op if the tunnel is already established or being established // Handshake will attempt to initiate a tunnel with the provided vpn address. This is a no-op if the tunnel is already established or being established
@@ -393,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) { 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) err := newPacket(p, false, fp)
if err != nil { if err != nil {
f.l.Warn("error while parsing outgoing packet for firewall check", "error", err) f.l.Warn("error while parsing outgoing packet for firewall check", "error", err)
@@ -401,7 +379,7 @@ func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubTyp
} }
// check if packet is in outbound fw rules // 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 dropReason != nil {
if f.l.Enabled(context.Background(), slog.LevelDebug) { if f.l.Enabled(context.Background(), slog.LevelDebug) {
f.l.Debug("dropping cached packet", f.l.Debug("dropping cached packet",
@@ -514,20 +492,14 @@ func (f *Interface) prepareSendVia(via *HostInfo,
// nb is a buffer used to store the nonce value, re-used for performance reasons. // nb is a buffer used to store the nonce value, re-used for performance reasons.
// out is a buffer used to store the result of the Encrypt operation // out is a buffer used to store the result of the Encrypt operation
// q indicates which writer to use to send the packet. // q indicates which writer to use to send the packet.
func (f *Interface) SendVia(via *HostInfo, func (f *Interface) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
relay *Relay,
ad,
nb,
out []byte,
nocopy bool,
) {
toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy) toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy)
if err != nil { if err != nil {
// already logged by prepareSendVia // already logged by prepareSendVia
return return
} }
err = f.writers[0].WriteTo(toSend, via.GetRemote()) err = f.writers[q].WriteTo(toSend, via.GetRemote())
if err != nil { if err != nil {
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err) via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
} }
@@ -601,7 +573,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
if err != nil { if err != nil {
hostinfo.logger(f.l).Error("Failed to write outgoing packet", hostinfo.logger(f.l).Error("Failed to write outgoing packet",
"error", err, "error", err,
"udpAddr", remote, "udpAddr", hr,
) )
} }
} else { } else {
@@ -616,7 +588,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
) )
continue continue
} }
f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true) f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true, q)
break break
} }
} }
+76 -66
View File
@@ -97,11 +97,6 @@ type Interface struct {
// a CPU at all (tun.pin_threads, default true). When false, threads are // a CPU at all (tun.pin_threads, default true). When false, threads are
// left free to migrate as on stock nebula. // left free to migrate as on stock nebula.
pinThreads bool pinThreads bool
// ecnEnabled gates RFC 6040 underlay ECN propagation. When true,
// inside.go copies the inner ECN onto the outer carrier on encap and
// decryptToTun folds outer CE into the inner header on decap. Toggle
// via tunnels.ecn (default true).
ecnEnabled atomic.Bool
relayManager *relayManager relayManager *relayManager
tryPromoteEvery atomic.Uint32 tryPromoteEvery atomic.Uint32
@@ -120,10 +115,12 @@ type Interface struct {
ctx context.Context ctx context.Context
writers []udp.Conn writers []udp.Conn
queues []tio.Queue queues []tio.Queue
// batchers is one per tun queue, wrapping queues[i]. // batchers is one per tun queue, wrapping queues[i]. readOutsidePackets
// decryptToTun sends plaintext into the batch.RxBatcher; // commits plaintext into the batcher; the plaintext is decrypted
// listenOut calls its Flush at the end of each UDP recvmmsg batch. // in place inside the UDP receive buffers, so listenOut must call Flush
batchers []batch.RxBatcher // at the end of each UDP recvmmsg batch, before those buffers are
// reused (every udp.Conn ListenOut guarantees that ordering).
batchers []*batch.MultiCoalescer
wg sync.WaitGroup wg sync.WaitGroup
// fatalErr holds the first unexpected reader error that caused shutdown. // fatalErr holds the first unexpected reader error that caused shutdown.
@@ -135,18 +132,13 @@ type Interface struct {
metricHandshakes metrics.Histogram metricHandshakes metrics.Histogram
messageMetrics *MessageMetrics messageMetrics *MessageMetrics
cachedPacketMetrics *cachedPacketMetrics cachedPacketMetrics *cachedPacketMetrics
metricTxDropped metrics.Counter
l *slog.Logger l *slog.Logger
} }
type EncWriter interface { type EncWriter interface {
SendVia(via *HostInfo, SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int)
relay *Relay,
ad,
nb,
out []byte,
nocopy bool,
)
SendMessageToVpnAddr(t header.MessageType, st header.MessageSubType, vpnAddr netip.Addr, p, nb, out []byte) SendMessageToVpnAddr(t header.MessageType, st header.MessageSubType, vpnAddr netip.Addr, p, nb, out []byte)
SendMessageToHostInfo(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) SendMessageToHostInfo(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte)
Handshake(vpnAddr netip.Addr) Handshake(vpnAddr netip.Addr)
@@ -205,6 +197,10 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
return nil, errors.New("no connection manager") return nil, errors.New("no connection manager")
} }
if c.routines <= 1 {
c.PinThreads = false //pinning is not useful unless there's more than one tun reader
}
cs := c.pki.getCertState() cs := c.pki.getCertState()
ifce := &Interface{ ifce := &Interface{
ctx: ctx, ctx: ctx,
@@ -222,7 +218,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
routines: c.routines, routines: c.routines,
version: c.version, version: c.version,
writers: make([]udp.Conn, c.routines), writers: make([]udp.Conn, c.routines),
batchers: make([]batch.RxBatcher, c.routines), batchers: make([]*batch.MultiCoalescer, c.routines),
myVpnNetworks: cs.myVpnNetworks, myVpnNetworks: cs.myVpnNetworks,
myVpnNetworksTable: cs.myVpnNetworksTable, myVpnNetworksTable: cs.myVpnNetworksTable,
myVpnAddrs: cs.myVpnAddrs, myVpnAddrs: cs.myVpnAddrs,
@@ -235,6 +231,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
pinThreads: c.PinThreads, pinThreads: c.PinThreads,
metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)), metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)),
metricTxDropped: metrics.GetOrRegisterCounter("udp.tx.dropped", nil),
messageMetrics: c.MessageMetrics, messageMetrics: c.MessageMetrics,
cachedPacketMetrics: &cachedPacketMetrics{ cachedPacketMetrics: &cachedPacketMetrics{
sent: metrics.GetOrRegisterCounter("hostinfo.cached_packets.sent", nil), sent: metrics.GetOrRegisterCounter("hostinfo.cached_packets.sent", nil),
@@ -288,6 +285,13 @@ func (f *Interface) activate() error {
return err return err
} }
if len(queues) < f.routines { if len(queues) < f.routines {
// TODO: this clamp is only safe because it is unreachable when the
// udp side has multiple readers (linux Queues opens exactly n or
// errors; every other platform already clamped routines to 1 above).
// If a platform ever returns fewer queues than routines with
// SO_REUSEPORT sockets already bound, the surplus sockets get no
// listenOut and the kernel blackholes every flow it hashes to them —
// fail loudly or close the extra sockets instead.
f.l.Warn("tun multiqueue is not supported on this platform, falling back to fewer routines", f.l.Warn("tun multiqueue is not supported on this platform, falling back to fewer routines",
"requested", f.routines, "opened", len(queues)) "requested", f.routines, "opened", len(queues))
f.routines = len(queues) f.routines = len(queues)
@@ -297,18 +301,7 @@ func (f *Interface) activate() error {
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines)) metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
for i := range f.queues { for i := range f.queues {
caps := tio.QueueCapabilities(f.queues[i]) f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l)
if caps.TSO || caps.USO {
// Multi-lane: TCP gets coalesced when TSO is on, UDP when USO
// is on, everything else (and either lane disabled) falls
// through to passthrough so non-IP / non-TCP-UDP traffic still
// reaches the TUN.
arena := batch.NewArena(batch.DefaultMultiArenaCap)
f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l, arena, caps.TSO, caps.USO)
} else {
arena := batch.NewArena(batch.DefaultPassthroughArenaCap)
f.batchers[i] = batch.NewPassthrough(f.queues[i], arena.Reserve, arena.Reset)
}
} }
// On error the caller owns the cleanup, Control.Start cancels the service context // On error the caller owns the cleanup, Control.Start cancels the service context
@@ -356,6 +349,31 @@ func (f *Interface) onFatal(err error) {
} }
} }
type rxContext struct {
q int
scratch []byte
// nb is a re-usable nonce buffer for decrypt calls to use
nb []byte
h *header.H
fwPacket *firewall.ParsedPacket
hostmapCache map[uint32]*HostInfo
lhh *LightHouseHandler
ctCache *firewall.ConntrackCacheTicker
}
func newRxContext(f *Interface, q int) *rxContext {
return &rxContext{
q: q,
scratch: make([]byte, mtu),
nb: make([]byte, 12, 12),
h: &header.H{},
fwPacket: &firewall.ParsedPacket{},
hostmapCache: map[uint32]*HostInfo{},
lhh: f.lightHouse.NewRequestHandler(),
ctCache: firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout),
}
}
func (f *Interface) listenOut(i int) { func (f *Interface) listenOut(i int) {
var li udp.Conn var li udp.Conn
if i > 0 { if i > 0 {
@@ -364,21 +382,17 @@ func (f *Interface) listenOut(i int) {
li = f.outside li = f.outside
} }
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout) rxc := newRxContext(f, i)
lhh := f.lightHouse.NewRequestHandler()
h := &header.H{}
fwPacket := &firewall.Packet{}
nb := make([]byte, 12, 12)
listener := func(fromUdpAddr netip.AddrPort, payload []byte, meta udp.RxMeta) { listener := func(fromUdpAddr netip.AddrPort, payload []byte) {
plaintext := f.batchers[i].Reserve(len(payload)) f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, payload, rxc)
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get(), meta)
} }
flusher := func() { flusher := func() {
if err := f.batchers[i].Flush(); err != nil { if err := f.batchers[i].Flush(); err != nil {
f.l.Error("Failed to flush tun coalescer", "error", err) f.l.Error("Failed to flush tun coalescer", "error", err)
} }
clear(rxc.hostmapCache)
} }
err := li.ListenOut(listener, flusher) err := li.ListenOut(listener, flusher)
@@ -394,10 +408,7 @@ func (f *Interface) listenOut(i int) {
f.l.Debug("underlay reader is done", "reader", i) f.l.Debug("underlay reader is done", "reader", i)
} }
func (f *Interface) listenIn(queue tio.Queue, i int) { func (f *Interface) pinThisThread(i int) {
// Pinning this thread (and goroutine) to a single CPU keeps every sendmmsg from this goroutine going through the
// same TX ring on the nic, so the wire sees per-flow order. Skip entirely when tun.pin_threads is false.
if f.pinThreads {
var cpu int var cpu int
if n := len(f.cpuAffinity); n > 0 { if n := len(f.cpuAffinity); n > 0 {
// Explicit tun.cpu_affinity list wins; parseCpuAffinity already // Explicit tun.cpu_affinity list wins; parseCpuAffinity already
@@ -414,12 +425,19 @@ func (f *Interface) listenIn(queue tio.Queue, i int) {
if err := util.PinThreadToCPU(cpu); err != nil { if err := util.PinThreadToCPU(cpu); err != nil {
f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err) f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err)
} }
}
func (f *Interface) listenIn(queue tio.Queue, i int) {
// Pinning this thread (and goroutine) to a single CPU keeps every sendmmsg from this goroutine going through the
// same TX ring on the nic, so the wire sees per-flow order. Skip entirely when tun.pin_threads is false.
if f.pinThreads {
f.pinThisThread(i)
} }
rejectBuf := make([]byte, mtu) rejectBuf := make([]byte, mtu)
arenaSize := batch.SendBatchCap * (udp.MTU + 32) arenaSize := batch.SendBatchCap * (udp.MTU + 32)
sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, arenaSize) sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, arenaSize)
fwPacket := &firewall.Packet{} fwPacket := &firewall.ParsedPacket{}
nb := make([]byte, 12, 12) nb := make([]byte, 12, 12)
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout) conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
@@ -441,26 +459,35 @@ func (f *Interface) listenIn(queue tio.Queue, i int) {
// accumulated so the first packets of a deep read drain // accumulated so the first packets of a deep read drain
// hit the wire while the rest are still being encrypted. // hit the wire while the rest are still being encrypted.
if sb.Len() >= batch.SendBatchCap { if sb.Len() >= batch.SendBatchCap {
if err := sb.Flush(); err != nil { f.flushSendBatch(sb, i)
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
} }
} }
} f.flushSendBatch(sb, i)
if err := sb.Flush(); err != nil {
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
}
} }
f.l.Debug("overlay reader is done", "reader", i) f.l.Debug("overlay reader is done", "reader", i)
} }
// flushSendBatch drains sb to the underlay and accounts for anything it could not deliver. A shortfall means
// specific destinations were undeliverable (a stale remote, a reject rule), which the backend logs per peer at
// debug; here it is only a counter, so one unreachable peer cannot spam a log line per batch.
func (f *Interface) flushSendBatch(sb *batch.SendBatch, q int) {
queued := sb.Len()
written, err := sb.Flush()
if err != nil {
f.l.Error("Failed to write outgoing batch", "error", err, "writer", q)
}
if dropped := queued - written; dropped > 0 {
f.metricTxDropped.Inc(int64(dropped))
}
}
func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) { func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
c.RegisterReloadCallback(f.reloadFirewall) c.RegisterReloadCallback(f.reloadFirewall)
c.RegisterReloadCallback(f.reloadSendRecvError) c.RegisterReloadCallback(f.reloadSendRecvError)
c.RegisterReloadCallback(f.reloadAcceptRecvError) c.RegisterReloadCallback(f.reloadAcceptRecvError)
c.RegisterReloadCallback(f.reloadDisconnectInvalid) c.RegisterReloadCallback(f.reloadDisconnectInvalid)
c.RegisterReloadCallback(f.reloadMisc) c.RegisterReloadCallback(f.reloadMisc)
c.RegisterReloadCallback(f.reloadEcn)
for _, udpConn := range f.writers { for _, udpConn := range f.writers {
c.RegisterReloadCallback(udpConn.ReloadConfig) c.RegisterReloadCallback(udpConn.ReloadConfig)
@@ -593,23 +620,6 @@ func (f *Interface) reloadMisc(c *config.C) {
} }
} }
// reloadEcn syncs Interface.ecnEnabled with the tunnels.ecn config knob.
// Default is enabled (RFC 6040 normal mode); set false on the rare path
// where an underlay middlebox rewrites or drops ECN bits unpredictably.
func (f *Interface) reloadEcn(c *config.C) {
initial := c.InitialLoad()
if initial || c.HasChanged("tunnels.ecn") {
v := c.GetBool("tunnels.ecn", true)
changed := f.ecnEnabled.Swap(v) != v
if !initial {
f.l.Info("tunnels.ecn changed", "enabled", v)
if changed {
f.l.Warn("tunnels.ecn datapath toggled, but route-level ECN negotiation (RTAX_FEATURE_ECN) retains its previous state until nebula is restarted", "enabled", v)
}
}
}
}
func (f *Interface) emitStats(ctx context.Context, i time.Duration) { func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
ticker := time.NewTicker(i) ticker := time.NewTicker(i)
defer ticker.Stop() defer ticker.Stop()
+27 -3
View File
@@ -199,7 +199,7 @@ func ipv4CreateRejectTCPPacket(packet []byte, out []byte) []byte {
} }
func ipv6CreateRejectPacket(packet []byte, out []byte) []byte { func ipv6CreateRejectPacket(packet []byte, out []byte) []byte {
proto, offset, isFragment := ipv6FindUpperProtocol(packet) proto, offset, isFragment := IPv6FindUpperProtocol(packet)
if isFragment { if isFragment {
return nil return nil
} }
@@ -333,11 +333,34 @@ func ipv6CreateRejectTCPPacket(packet []byte, out []byte, offset int) []byte {
return out return out
} }
func ipv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool) { // maxIPv6ExtHeaders caps the extension-header walk in IPv6FindUpperProtocol.
// RFC 8200 legal chains are shorter (each header at most once, Destination
// Options at most twice), so the cap only bites crafted packets, which would
// otherwise make us walk their whole payload 8 bytes at a time.
const maxIPv6ExtHeaders = 8
// IPv6FindUpperProtocol walks packet's IPv6 extension-header chain and
// returns the terminal (upper-layer) protocol number, the byte offset where
// that protocol's header begins, and whether the packet is a non-first
// fragment. It steps over Hop-by-Hop (0), Routing (43), Fragment (44),
// AH (51), and Destination Options (60); anything else — including ESP,
// whose payload is encrypted — terminates the walk.
//
// For a non-first fragment, nextHeader still names the flow's upper
// protocol (copied from the fragment header) but offset points at fragment
// payload, not a real transport header: consult isFragment before
// dereferencing offset. If the chain is truncated, over-long, or the packet
// is shorter than an IPv6 header, the walk stops early and nextHeader is
// whatever it stopped on (59, IPPROTO_NONE, for the too-short case) —
// callers treat any non-transport result as unclassifiable.
func IPv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool) {
if len(packet) < ipv6.HeaderLen {
return 59, 0, false // IPPROTO_NONE: nothing to classify
}
nextHeader = packet[6] nextHeader = packet[6]
offset = ipv6.HeaderLen offset = ipv6.HeaderLen
for { for range maxIPv6ExtHeaders {
switch nextHeader { switch nextHeader {
case 0, 43, 60: // Hop-by-Hop, Routing, Destination case 0, 43, 60: // Hop-by-Hop, Routing, Destination
if len(packet) < offset+2 { if len(packet) < offset+2 {
@@ -367,6 +390,7 @@ func ipv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragm
return nextHeader, offset, isFragment return nextHeader, offset, isFragment
} }
} }
return nextHeader, offset, isFragment
} }
func CreateICMPEchoResponse(packet, out []byte) []byte { func CreateICMPEchoResponse(packet, out []byte) []byte {
+118
View File
@@ -515,3 +515,121 @@ func TestCreateICMPEchoResponse_IPv6_NotICMPv6(t *testing.T) {
result := CreateICMPEchoResponse(packet, out) result := CreateICMPEchoResponse(packet, out)
assert.Nil(t, result) assert.Nil(t, result)
} }
func TestIPv6FindUpperProtocol(t *testing.T) {
src := net.ParseIP("fd00::1")
dst := net.ParseIP("fd00::2")
// extHdr builds one 8-byte-unit extension header: next, hdrExtLen
// ((extra+1)*8 bytes total), padded to size.
extHdr := func(next uint8, extra int) []byte {
b := make([]byte, (extra+1)*8)
b[0] = next
b[1] = uint8(extra)
return b
}
t.Run("no extension headers", func(t *testing.T) {
for _, proto := range []uint8{6, 17, 58} {
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, proto, make([]byte, 20)))
assert.Equal(t, proto, nh)
assert.Equal(t, ipv6.HeaderLen, offset)
assert.False(t, frag)
}
})
t.Run("hop-by-hop then TCP", func(t *testing.T) {
payload := append(extHdr(6, 0), make([]byte, 20)...)
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, payload))
assert.Equal(t, uint8(6), nh)
assert.Equal(t, ipv6.HeaderLen+8, offset)
assert.False(t, frag)
})
t.Run("chained headers honor length units", func(t *testing.T) {
// Hop-by-Hop (8B) -> Dest Options (16B) -> Routing (8B) -> UDP.
payload := extHdr(60, 0)
payload = append(payload, extHdr(43, 1)...)
payload = append(payload, extHdr(17, 0)...)
payload = append(payload, make([]byte, 8)...)
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, payload))
assert.Equal(t, uint8(17), nh)
assert.Equal(t, ipv6.HeaderLen+8+16+8, offset)
assert.False(t, frag)
})
t.Run("AH length is in 4-byte units plus 2", func(t *testing.T) {
// AH payload-len byte 4 -> (4+2)*4 = 24 bytes on the wire.
ah := make([]byte, 24)
ah[0] = 6
ah[1] = 4
payload := append(ah, make([]byte, 20)...)
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 51, payload))
assert.Equal(t, uint8(6), nh)
assert.Equal(t, ipv6.HeaderLen+24, offset)
assert.False(t, frag)
})
t.Run("first fragment walks to the transport header", func(t *testing.T) {
frag := make([]byte, 8)
frag[0] = 17
binary.BigEndian.PutUint16(frag[2:4], 0x0001) // offset 0, M=1
payload := append(frag, make([]byte, 8)...)
nh, offset, isFrag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 44, payload))
assert.Equal(t, uint8(17), nh)
assert.Equal(t, ipv6.HeaderLen+8, offset)
assert.False(t, isFrag, "first fragment carries the real transport header")
})
t.Run("non-first fragment is flagged", func(t *testing.T) {
frag := make([]byte, 8)
frag[0] = 17
binary.BigEndian.PutUint16(frag[2:4], 1<<3) // offset 1, M=0
payload := append(frag, make([]byte, 8)...)
nh, _, isFrag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 44, payload))
assert.Equal(t, uint8(17), nh, "fragment header still names the flow's L4")
assert.True(t, isFrag, "offset points at fragment payload, not a header")
})
t.Run("ESP terminates the walk", func(t *testing.T) {
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 50, make([]byte, 16)))
assert.Equal(t, uint8(50), nh)
assert.Equal(t, ipv6.HeaderLen, offset)
assert.False(t, frag)
})
t.Run("unknown protocol terminates the walk", func(t *testing.T) {
nh, offset, _ := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 132, make([]byte, 16))) // SCTP
assert.Equal(t, uint8(132), nh)
assert.Equal(t, ipv6.HeaderLen, offset)
})
t.Run("truncated extension header stops the walk", func(t *testing.T) {
// Next header says Hop-by-Hop but the packet ends at the IPv6 header.
nh, offset, frag := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, nil))
assert.Equal(t, uint8(0), nh, "unresolvable chain returns the extension header it stopped on")
assert.Equal(t, ipv6.HeaderLen, offset)
assert.False(t, frag)
})
t.Run("crafted over-long chain hits the cap", func(t *testing.T) {
// Ten chained Hop-by-Hop headers, then TCP. Illegal per RFC 8200
// (Hop-by-Hop may only appear first); the cap must stop the walk
// before it resolves rather than crawling arbitrary crafted chains.
var payload []byte
for i := 0; i < 9; i++ {
payload = append(payload, extHdr(0, 0)...)
}
payload = append(payload, extHdr(6, 0)...)
payload = append(payload, make([]byte, 20)...)
nh, _, _ := IPv6FindUpperProtocol(makeIPv6Packet(src, dst, 0, payload))
assert.Equal(t, uint8(0), nh, "walk must stop at the cap, not resolve to TCP")
})
t.Run("packet shorter than an IPv6 header", func(t *testing.T) {
nh, offset, frag := IPv6FindUpperProtocol(make([]byte, 39))
assert.Equal(t, uint8(59), nh) // IPPROTO_NONE
assert.Equal(t, 0, offset)
assert.False(t, frag)
})
}
+9 -1
View File
@@ -36,6 +36,10 @@ type LightHouse struct {
myVpnNetworksTable *bart.Lite myVpnNetworksTable *bart.Lite
punchy *Punchy punchy *Punchy
// localAddrsFn enumerates the underlay addresses we advertise. It is a field so tests can supply simulated
// addresses rather than whatever this machine's NICs happen to be. Set it before Start.
localAddrsFn func(*LocalAllowList) []netip.Addr
// Local cache of answers from light houses // Local cache of answers from light houses
// map of vpn addr to answers // map of vpn addr to answers
addrMap map[netip.Addr]*RemoteList addrMap map[netip.Addr]*RemoteList
@@ -107,6 +111,10 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)), queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
l: l, l: l,
} }
h.localAddrsFn = func(al *LocalAllowList) []netip.Addr {
return localAddrs(h.l, al)
}
lighthouses := make([]netip.Addr, 0) lighthouses := make([]netip.Addr, 0)
h.lighthouses.Store(&lighthouses) h.lighthouses.Store(&lighthouses)
staticList := make(map[netip.Addr]struct{}) staticList := make(map[netip.Addr]struct{})
@@ -918,7 +926,7 @@ func (lh *LightHouse) SendUpdate() {
} }
lal := lh.GetLocalAllowList() lal := lh.GetLocalAllowList()
for _, e := range localAddrs(lh.l, lal) { for _, e := range lh.localAddrsFn(lal) {
if lh.myVpnNetworksTable.Contains(e) { if lh.myVpnNetworksTable.Contains(e) {
continue continue
} }
+1 -1
View File
@@ -498,7 +498,7 @@ type testEncWriter struct {
protocolVersion cert.Version protocolVersion cert.Version
} }
func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool) { func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool, q int) {
} }
func (tw *testEncWriter) Handshake(vpnIp netip.Addr) { func (tw *testEncWriter) Handshake(vpnIp netip.Addr) {
} }
+23 -55
View File
@@ -6,12 +6,14 @@ import (
"log/slog" "log/slog"
"net" "net"
"net/netip" "net/netip"
"os"
"runtime/debug" "runtime/debug"
"slices" "slices"
"strings" "strings"
"time" "time"
"github.com/slackhq/nebula/config" "github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/cpupick"
"github.com/slackhq/nebula/overlay" "github.com/slackhq/nebula/overlay"
"github.com/slackhq/nebula/sshd" "github.com/slackhq/nebula/sshd"
"github.com/slackhq/nebula/udp" "github.com/slackhq/nebula/udp"
@@ -165,7 +167,13 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
for i := 0; i < routines; i++ { for i := 0; i < routines; i++ {
l.Info("listening", "addr", netip.AddrPortFrom(listenHost, uint16(port))) l.Info("listening", "addr", netip.AddrPortFrom(listenHost, uint16(port)))
udpServer, err := udp.NewListener(l, listenHost, port, routines > 1, c.GetInt("listen.batch", 64)) batchSize := c.GetInt("listen.batch", 64)
if batchSize < 1 {
oldBatch := batchSize
batchSize = 1
l.Warn("listen.batch size is invalid", "provided", oldBatch, "overridden to", batchSize)
}
udpServer, err := udp.NewListener(l, listenHost, port, routines > 1, batchSize)
if err != nil { if err != nil {
return nil, util.NewContextualError("Failed to open udp listener", m{"queue": i}, err) return nil, util.NewContextualError("Failed to open udp listener", m{"queue": i}, err)
} }
@@ -216,8 +224,17 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
pinThreads := c.GetBool("tun.pin_threads", true) pinThreads := c.GetBool("tun.pin_threads", true)
cpuAffinity := parseCpuAffinity(c, l, routines) cpuAffinity := parseCpuAffinity(c, l, routines)
if pinThreads && len(cpuAffinity) == 0 && !configTest { if pinThreads && routines > 1 && len(cpuAffinity) == 0 && !configTest {
cpuAffinity = defaultCPUAffinityAvoidingIRQs(l, routines) // The operator didn't choose pin CPUs, so pick a default set that
// prefers performance cores and doesn't stack co-located instances
// onto allowed[0]. The bound UDP port keys the per-instance spread:
// distinct across instances sharing a box, stable across restarts.
// A nil result keeps listenIn's stock allowed[i] fallback.
key := uint64(os.Getpid())
if ap, err := udpConns[0].LocalAddr(); err == nil && ap.Port() != 0 {
key = uint64(ap.Port())
}
cpuAffinity = cpupick.Default(routines, key, l)
} }
ifConfig := &InterfaceConfig{ ifConfig := &InterfaceConfig{
@@ -260,7 +277,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
ifce.reloadDisconnectInvalid(c) ifce.reloadDisconnectInvalid(c)
ifce.reloadSendRecvError(c) ifce.reloadSendRecvError(c)
ifce.reloadAcceptRecvError(c) ifce.reloadAcceptRecvError(c)
ifce.reloadEcn(c)
handshakeManager.f = ifce handshakeManager.f = ifce
go handshakeManager.Run(ctx) go handshakeManager.Run(ctx)
@@ -281,6 +297,8 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
attachCommands(l, c, ssh, ifce) attachCommands(l, c, ssh, ifce)
networkChanges := udp.NewNetworkChangeMonitor(ctx, l, c)
return &Control{ return &Control{
state: StateReady, state: StateReady,
f: ifce, f: ifce,
@@ -291,6 +309,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
statsStart: stats.Start, statsStart: stats.Start,
dnsStart: ds.Start, dnsStart: ds.Start,
lighthouseStart: lightHouse.StartUpdateWorker, lighthouseStart: lightHouse.StartUpdateWorker,
networkChangeStart: networkChanges.Start,
connectionManagerStart: connManager.Start, connectionManagerStart: connManager.Start,
}, nil }, nil
} }
@@ -359,57 +378,6 @@ func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int {
return cpus return cpus
} }
// defaultCPUAffinityAvoidingIRQs picks the default pin set for the tun
// readers when tun.cpu_affinity is unset: allowed CPUs that do NOT service
// any physical NIC's interrupts. The stock allowed[i] spread pins the
// encrypt threads onto exactly the cores most drivers affine their first RX
// queue IRQs to, so whenever a flow's RSS queue fires on a core hosting a
// tun reader, NAPI and encrypt fight for the core and per-flow throughput
// drops (measured: REV 8.4 vs 10.2 Gbps on the same hardware, 2026-07-14).
//
// Returns nil — keeping the old allowed[i] fallback in listenIn — when IRQ
// info is unavailable or when there aren't enough IRQ-free CPUs to give
// every routine its own core: silently doubling readers up on fewer cores
// is worse than the occasional IRQ collision. NICs whose vectors blanket
// every CPU (e.g. mlx5 defaults to one queue per core) make avoidance
// impossible; narrowing the NIC's spread (ethtool -X <dev> equal N, or
// /proc/irq/*/smp_affinity) or setting tun.cpu_affinity explicitly makes it
// effective.
func defaultCPUAffinityAvoidingIRQs(l *slog.Logger, routines int) []int {
irq, err := util.NICIRQCPUs()
if err != nil || len(irq) == 0 {
return nil
}
allowed, err := util.AllowedCPUs()
if err != nil {
return nil
}
cpus := chooseIRQFreeCPUs(allowed, irq, routines)
if cpus == nil {
l.Info("not enough CPUs are free of NIC IRQs to give every tun reader its own; using the default spread",
"routines", routines, "allowed", len(allowed), "irqCPUs", len(irq))
return nil
}
l.Info("pinning tun readers to CPUs clear of NIC IRQs", "cpus", cpus)
return cpus
}
// chooseIRQFreeCPUs returns the first `routines` allowed CPUs not present in
// irq, or nil if fewer than `routines` qualify.
func chooseIRQFreeCPUs(allowed []int, irq map[int]bool, routines int) []int {
free := make([]int, 0, routines)
for _, cpu := range allowed {
if irq[cpu] {
continue
}
free = append(free, cpu)
if len(free) == routines {
return free
}
}
return nil
}
func moduleVersion() string { func moduleVersion() string {
info, ok := debug.ReadBuildInfo() info, ok := debug.ReadBuildInfo()
if !ok { if !ok {
-20
View File
@@ -9,26 +9,6 @@ import (
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
) )
func TestChooseIRQFreeCPUs(t *testing.T) {
irq := map[int]bool{0: true, 1: true, 2: true, 3: true}
// Plenty of IRQ-free CPUs: take the first `routines` of them in order.
assert.Equal(t, []int{4, 5}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 2))
// Exactly enough.
assert.Equal(t, []int{4, 5, 6}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 3))
// Not enough IRQ-free CPUs: nil, caller keeps the old default rather
// than doubling readers up on shared cores.
assert.Nil(t, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4}, irq, 2))
// No IRQ info at all behaves like a plain prefix of allowed.
assert.Equal(t, []int{0, 1}, chooseIRQFreeCPUs([]int{0, 1, 2}, map[int]bool{}, 2))
// Non-contiguous allowed set (cgroup cpuset) with holes.
assert.Equal(t, []int{9, 12}, chooseIRQFreeCPUs([]int{1, 3, 9, 12}, map[int]bool{1: true, 3: true}, 2))
}
func TestParseCpuAffinity(t *testing.T) { func TestParseCpuAffinity(t *testing.T) {
l := test.NewLogger() l := test.NewLogger()
+45
View File
@@ -164,3 +164,48 @@ func TestCipherStateNilSafety(t *testing.T) {
assert.Empty(t, out) assert.Empty(t, out)
assert.Equal(t, 0, cc.Overhead()) assert.Equal(t, 0, cc.Overhead())
} }
func TestCipherStateAESGCMInPlaceDecrypt(t *testing.T) {
enc, dec := buildCipherStates(t, CipherAESGCM)
inPlaceDecrypt(t, NewCipherStateAESGCM(enc), NewCipherStateAESGCM(dec))
}
func TestCipherStateChaChaPolyInPlaceDecrypt(t *testing.T) {
enc, dec := buildCipherStates(t, noise.CipherChaChaPoly)
inPlaceDecrypt(t, NewCipherStateChaChaPoly(enc), NewCipherStateChaChaPoly(dec))
}
func inPlaceDecrypt(t *testing.T, enc, dec CipherState) {
t.Helper()
const hdrLen = 16
plaintext := []byte("in-place decrypt should replace the ciphertext bytes")
nb := make([]byte, 12)
// packet = [16-byte header | ciphertext+tag], like a nebula Message.
packet := make([]byte, hdrLen, hdrLen+len(plaintext)+enc.Overhead())
for i := range packet {
packet[i] = byte(i)
}
packet, err := enc.EncryptDanger(packet, packet[:hdrLen], plaintext, 1, nb)
require.NoError(t, err)
// Simulate a GRO row: [packet | next segment]. A failed auth on packet
// may zero packet's plaintext region but must not touch the header, the
// tag, or the neighboring segment.
neighbor := []byte("next coalesced segment, must stay intact")
row := append(append([]byte(nil), packet...), neighbor...)
tampered := row[:len(packet)]
tampered[hdrLen] ^= 0x01
_, err = dec.DecryptDanger(tampered[hdrLen:hdrLen], tampered[:hdrLen], tampered[hdrLen:], 1, nb)
require.Error(t, err)
assert.Equal(t, packet[:hdrLen], tampered[:hdrLen], "failed auth must not touch the header")
assert.Equal(t, packet[len(packet)-dec.Overhead():], tampered[len(tampered)-dec.Overhead():],
"failed auth must not touch the tag")
assert.Equal(t, neighbor, row[len(packet):], "failed auth must not touch the next segment")
out, err := dec.DecryptDanger(packet[hdrLen:hdrLen], packet[:hdrLen], packet[hdrLen:], 1, nb)
require.NoError(t, err)
assert.Equal(t, plaintext, out)
// The plaintext must be IN the packet buffer, not a fresh allocation.
assert.Equal(t, &packet[hdrLen], &out[0], "plaintext must alias the packet buffer")
}
+78 -157
View File
@@ -13,7 +13,7 @@ import (
"github.com/slackhq/nebula/firewall" "github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/header" "github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/udp" "github.com/slackhq/nebula/overlay/batch"
"golang.org/x/net/ipv4" "golang.org/x/net/ipv4"
) )
@@ -23,7 +23,11 @@ const (
var ErrOutOfWindow = errors.New("out of window packet") var ErrOutOfWindow = errors.New("out of window packet")
func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) { // readOutsidePackets processes one received underlay packet.
// Message payloads are decrypted IN PLACE, so packet must stay untouched
// by the caller until the batcher for queue q has been flushed
func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, rxc *rxContext) {
h := rxc.h
err := h.Parse(packet) err := h.Parse(packet)
if err != nil { if err != nil {
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors // Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
@@ -91,7 +95,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
if isMessageRelay { if isMessageRelay {
hostinfo = f.hostMap.QueryRelayIndex(h.RemoteIndex) hostinfo = f.hostMap.QueryRelayIndex(h.RemoteIndex)
} else { } else {
hostinfo = f.hostMap.QueryIndex(h.RemoteIndex) hostinfo = f.hostMap.QueryIndexCached(h.RemoteIndex, rxc.hostmapCache)
} }
// At this point we should have a valid existing tunnel, verify and send // At this point we should have a valid existing tunnel, verify and send
@@ -103,26 +107,32 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
return return
} }
if len(packet) < header.Len+hostinfo.ConnectionState.dKey.Overhead() {
f.messageMetrics.RxInvalid(1)
if f.l.Enabled(context.Background(), slog.LevelDebug) {
f.l.Debug("packet too small", "from", via, "length", len(packet))
}
return
}
// All remaining packets are encrypted // All remaining packets are encrypted
ci := hostinfo.ConnectionState
if !ci.window.Check(f.l, h.MessageCounter) {
return
}
// Relay packets are special
if isMessageRelay { if isMessageRelay {
f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, lhf, nb, q, localCache, meta) // Relay packets are special, this branch should always early-return
return err = hostinfo.ConnectionState.VerifyRelay(f.l, h.MessageCounter, packet, rxc.nb)
}
out, err = f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
if err != nil { if err != nil {
if f.l.Enabled(context.Background(), slog.LevelDebug) { if f.l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(f.l).Debug("Failed to decrypt packet", hostinfo.logger(f.l).Debug("Failed to verify relay packet", "error", err, "from", via, "header", h)
"error", err, }
"from", via, return
"header", h, }
) f.handleOutsideRelayPacket(hostinfo, via, packet, rxc)
return
}
out, err := hostinfo.ConnectionState.Decrypt(f.l, h.MessageCounter, packet, rxc.nb)
if err != nil {
if f.l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(f.l).Debug("Failed to decrypt packet", "error", err, "from", via, "header", h)
} }
return return
} }
@@ -135,7 +145,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
case header.Message: case header.Message:
switch h.Subtype { switch h.Subtype {
case header.MessageNone: case header.MessageNone:
f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, nb, q, localCache, meta) f.handleOutsideMessagePacket(hostinfo, h.MessageCounter, out, rxc)
default: default:
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h) hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
return return
@@ -143,15 +153,23 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
case header.LightHouse: case header.LightHouse:
//TODO: assert via is not relayed //TODO: assert via is not relayed
lhf.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f) rxc.lhh.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f)
case header.Test: case header.Test:
switch h.Subtype { switch h.Subtype {
case header.TestReply: case header.TestReply:
// No-op, useful for the Roaming and connectionManager side-effects above // No-op, useful for the Roaming and connectionManager side-effects above
case header.TestRequest: case header.TestRequest:
//recycle the input packet ciphertext as our output buffer const maxCipherOverhead = 16 //todo we use this too often, needs a real importable const
f.send(header.Test, header.TestReply, ci, hostinfo, out, nb, packet) const maxOverhead = header.Len + header.Len + maxCipherOverhead + maxCipherOverhead
if maxOverhead+len(out) > len(rxc.scratch) {
// A reply that cannot fit in scratch is dropped no matter the log level.
if f.l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(f.l).Debug("dropping oversized test request", "payloadLen", len(out), "from", via)
}
return
}
f.send(header.Test, header.TestReply, hostinfo.ConnectionState, hostinfo, out, rxc.nb, rxc.scratch[:0])
default: default:
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h) hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h)
return return
@@ -169,28 +187,10 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
} }
} }
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) { func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, packet []byte, rxc *rxContext) {
// The entire body is sent as AD, not encrypted. h := rxc.h
// The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value. // Successfully validated the thing. Get rid of the Relay header and the AEAD tag
// The packet is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's signedPayload := packet[header.Len : len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
// otherwise malformed (meaning, there is no trailing 16 byte AEAD value), then this will result in at worst a 0-length slice
// which will gracefully fail in the DecryptDanger call.
signedPayload := packet[:len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
signatureValue := packet[len(packet)-hostinfo.ConnectionState.dKey.Overhead():]
var err error
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, signedPayload, signatureValue, h.MessageCounter, nb)
if err != nil {
return
}
// Advance the replay window now that the frame is authenticated
if !hostinfo.ConnectionState.window.Update(f.l, h.MessageCounter) {
if f.l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(f.l).Debug("dropping out of window relay packet", "header", h)
}
return
}
// Successfully validated the thing. Get rid of the Relay header.
signedPayload = signedPayload[header.Len:]
// Pull the Roaming parts up here, and return in all call paths. // Pull the Roaming parts up here, and return in all call paths.
f.handleHostRoaming(hostinfo, via) f.handleHostRoaming(hostinfo, via)
// Track usage of both the HostInfo and the Relay for the received & authenticated packet // Track usage of both the HostInfo and the Relay for the received & authenticated packet
@@ -201,9 +201,7 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
if !ok { if !ok {
// The only way this happens is if hostmap has an index to the correct HostInfo, but the HostInfo is missing // The only way this happens is if hostmap has an index to the correct HostInfo, but the HostInfo is missing
// its internal mapping. This should never happen. // its internal mapping. This should never happen.
hostinfo.logger(f.l).Error("HostInfo missing remote relay index", hostinfo.logger(f.l).Error("HostInfo missing remote relay index", "relayRemoteIndex", h.RemoteIndex)
"relayRemoteIndex", h.RemoteIndex,
)
return return
} }
@@ -214,11 +212,10 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
via = ViaSender{ via = ViaSender{
UdpAddr: via.UdpAddr, UdpAddr: via.UdpAddr,
relayHI: hostinfo, relayHI: hostinfo,
remoteIdx: relay.RemoteIndex,
relay: relay, relay: relay,
IsRelayed: true, IsRelayed: true,
} }
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache, meta) f.readOutsidePackets(via, signedPayload, rxc)
case ForwardingType: case ForwardingType:
// Find the target HostInfo relay object // Find the target HostInfo relay object
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr) targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
@@ -235,9 +232,11 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
if targetRelay.State == Established { if targetRelay.State == Established {
switch targetRelay.Type { switch targetRelay.Type {
case ForwardingType: case ForwardingType:
// Forward this packet through the relay tunnel // Forward this packet through the relay tunnel, rebuilding it in place.
// Find the target HostInfo //todo it would potentially be nice to batch these // Encode overwrites the old outer header, and the new AEAD tag lands where the old one was
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false) fwdBuf := packet[:0]
//todo it would potentially be nice to batch these
f.SendVia(targetHI, targetRelay, signedPayload, rxc.nb, fwdBuf, true, rxc.q)
case TerminalType: case TerminalType:
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal") hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
return return
@@ -318,7 +317,11 @@ var (
) )
// newPacket validates and parses the interesting bits for the firewall out of the ip and sub protocol headers // 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 so an early-error return cannot
// leak the previous packet's offsets.
fp.IPHdrLen = 0
fp.FragAny = false
if len(data) < 1 { if len(data) < 1 {
return ErrPacketTooShort return ErrPacketTooShort
} }
@@ -333,7 +336,7 @@ func newPacket(data []byte, incoming bool, fp *firewall.Packet) error {
return ErrUnknownIPVersion 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) dataLen := len(data)
if dataLen < ipv6.HeaderLen { if dataLen < ipv6.HeaderLen {
return ErrIPv6PacketTooShort return ErrIPv6PacketTooShort
@@ -359,6 +362,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
switch proto { switch proto {
case layers.IPProtocolESP, layers.IPProtocolNoNextHeader: case layers.IPProtocolESP, layers.IPProtocolNoNextHeader:
fp.Protocol = uint8(proto) fp.Protocol = uint8(proto)
fp.IPHdrLen = offset
fp.RemotePort = 0 fp.RemotePort = 0
fp.LocalPort = 0 fp.LocalPort = 0
fp.Fragment = false fp.Fragment = false
@@ -369,6 +373,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
return ErrIPv6PacketTooShort return ErrIPv6PacketTooShort
} }
fp.Protocol = uint8(proto) fp.Protocol = uint8(proto)
fp.IPHdrLen = offset
fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6 fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6
icmptype := data[offset+1] icmptype := data[offset+1]
switch icmptype { switch icmptype {
@@ -386,6 +391,9 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
} }
fp.Protocol = uint8(proto) 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 { if incoming {
fp.RemotePort = binary.BigEndian.Uint16(data[offset : offset+2]) fp.RemotePort = binary.BigEndian.Uint16(data[offset : offset+2])
fp.LocalPort = binary.BigEndian.Uint16(data[offset+2 : offset+4]) fp.LocalPort = binary.BigEndian.Uint16(data[offset+2 : offset+4])
@@ -403,6 +411,9 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
return ErrIPv6PacketTooShort return ErrIPv6PacketTooShort
} }
// A fragment shape the coalescer must not touch either way, first fragment included.
fp.FragAny = true
// Check if this is the first fragment // 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 fragmentOffset := binary.BigEndian.Uint16(data[offset+2:offset+4]) &^ uint16(0x7) // Remove the reserved and M flag bits
if fragmentOffset != 0 { if fragmentOffset != 0 {
@@ -444,7 +455,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
return ErrIPv6CouldNotFindPayload 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? // Do we at least have an ipv4 header worth of data?
if len(data) < ipv4.HeaderLen { if len(data) < ipv4.HeaderLen {
return ErrIPv4PacketTooShort return ErrIPv4PacketTooShort
@@ -461,6 +472,10 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
// Check if this is the second or further fragment of a fragmented packet. // Check if this is the second or further fragment of a fragmented packet.
flagsfrags := binary.BigEndian.Uint16(data[6:8]) flagsfrags := binary.BigEndian.Uint16(data[6:8])
fp.Fragment = (flagsfrags & 0x1FFF) != 0 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 // Firewall handles protocol checks
fp.Protocol = data[9] fp.Protocol = data[9]
@@ -504,117 +519,23 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
return nil return nil
} }
func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, out []byte, packet []byte, h *header.H, nb []byte) ([]byte, error) { func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, messageCounter uint64, out []byte, rxc *rxContext) {
var err error err := newPacket(out, true, rxc.fwPacket)
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], mc, nb)
if err != nil { if err != nil {
return nil, err hostinfo.logger(f.l).Warn("Error while validating inbound packet", "error", err, "packet", out)
}
if !hostinfo.ConnectionState.window.Update(f.l, mc) {
return nil, ErrOutOfWindow
}
return out, nil
}
// 2-bit IP-level ECN codepoints (lower bits of IPv4 ToS / IPv6 TC).
const (
ecnNotECT = 0x00
ecnECT1 = 0x01
ecnECT0 = 0x02
ecnCE = 0x03
)
// applyOuterECN folds an outer CE mark from the underlay into the inner
// IP header per RFC 6040 normal mode. It mutates pkt[1] in place. Other
// codepoints are advisory only and leave the inner unchanged.
//
// Merge cases (outer × inner → action):
//
// outer != CE : no-op (inner is authoritative)
// outer == CE, inner Not-ECT : log; cannot propagate to a non-ECN host
// outer == CE, inner ECT/CE : rewrite inner ECN to CE
func applyOuterECN(pkt []byte, outerECN byte, hostinfo *HostInfo, l *slog.Logger) {
if outerECN&ecnCE != ecnCE || len(pkt) < 2 {
return
}
switch pkt[0] >> 4 {
case 4:
switch pkt[1] & 0x03 {
case ecnNotECT:
if l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
}
case ecnCE:
// Already CE.
default:
// Rewriting the ToS byte invalidates the IPv4 header checksum, so
// patch it incrementally per RFC 1624 (HC' = ~(~HC + ~m + m')). The
// ToS is the low byte of the 16-bit word at pkt[0:2]; the header
// checksum lives at pkt[10:12]. A header too short to carry a
// checksum can't be fixed up here, so leave it for newPacket to
// reject rather than emit a mangled packet.
if len(pkt) < ipv4.HeaderLen {
return
}
m := binary.BigEndian.Uint16(pkt[0:2])
pkt[1] = (pkt[1] &^ 0x03) | ecnCE
mNew := binary.BigEndian.Uint16(pkt[0:2])
sum := uint32(^binary.BigEndian.Uint16(pkt[10:12])) + uint32(^m) + uint32(mNew)
for sum > 0xffff {
sum = (sum >> 16) + (sum & 0xffff)
}
binary.BigEndian.PutUint16(pkt[10:12], ^uint16(sum))
}
case 6:
switch (pkt[1] >> 4) & 0x03 {
case ecnNotECT:
if l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
}
case ecnCE:
// Already CE.
default:
pkt[1] = (pkt[1] &^ 0x30) | (ecnCE << 4)
}
}
}
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
// RFC 6040 normal-mode combine: fold any outer CE mark stamped by the
// underlay into the inner header before firewall + TUN write. Other
// outer codepoints are advisory only — we keep the inner unchanged.
if f.ecnEnabled.Load() {
applyOuterECN(out, meta.OuterECN, hostinfo, f.l)
}
err := newPacket(out, true, fwPacket)
if err != nil {
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
"error", err,
"packet", out,
)
return return
} }
dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache) dropReason := f.firewall.Drop(rxc.fwPacket.Packet, true, hostinfo, f.pki.GetCAPool(), rxc.ctCache.Get())
if dropReason != nil { if dropReason != nil {
// NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, rxc.nb, rxc.scratch, rxc.q)
// This gives us a buffer to build the reject packet in. With UDP GRO this is a single segment of a shared
// recvmmsg row whose capacity runs to the end of the whole row, so cap it to its own length (cap==len) to
// keep the reject builder from writing past this segment into the next, not-yet-processed coalesced segment.
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet[:len(packet):len(packet)], q)
if f.l.Enabled(context.Background(), slog.LevelDebug) { if f.l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(f.l).Debug("dropping inbound packet", hostinfo.logger(f.l).Debug("dropping inbound packet", "fwPacket", rxc.fwPacket, "reason", dropReason)
"fwPacket", fwPacket,
"reason", dropReason,
)
} }
return return
} }
err = f.batchers[q].Commit(out) err = f.batchers[rxc.q].Commit(out, batch.SortKey{Epoch: hostinfo.ConnectionState.epoch, Counter: messageCounter}, rxc.fwPacket)
if err != nil { if err != nil {
f.l.Error("Failed to write to tun", "error", err) f.l.Error("Failed to write to tun", "error", err)
} }
+90 -5
View File
@@ -17,7 +17,7 @@ import (
) )
func Test_newPacket(t *testing.T) { func Test_newPacket(t *testing.T) {
p := &firewall.Packet{} p := &firewall.ParsedPacket{}
// length fails // length fails
err := newPacket([]byte{}, true, p) err := newPacket([]byte{}, true, p)
@@ -96,7 +96,7 @@ func Test_newPacket(t *testing.T) {
} }
func Test_newPacket_v6(t *testing.T) { func Test_newPacket_v6(t *testing.T) {
p := &firewall.Packet{} p := &firewall.ParsedPacket{}
// invalid ipv6 // invalid ipv6
ip := layers.IPv6{ ip := layers.IPv6{
@@ -345,7 +345,7 @@ func Test_newPacket_v6(t *testing.T) {
} }
func Test_newPacket_ipv6Fragment(t *testing.T) { func Test_newPacket_ipv6Fragment(t *testing.T) {
p := &firewall.Packet{} p := &firewall.ParsedPacket{}
ip := &layers.IPv6{ ip := &layers.IPv6{
Version: 6, Version: 6,
@@ -525,7 +525,7 @@ func BenchmarkParseV6(b *testing.B) {
secondFrag = append(secondFrag, fragHeader...) secondFrag = append(secondFrag, fragHeader...)
secondFrag = append(secondFrag, []byte{0xde, 0xad, 0xbe, 0xef}...) secondFrag = append(secondFrag, []byte{0xde, 0xad, 0xbe, 0xef}...)
fp := &firewall.Packet{} fp := &firewall.ParsedPacket{}
b.Run("Normal", func(b *testing.B) { b.Run("Normal", func(b *testing.B) {
for i := 0; i < b.N; i++ { 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 // host OS parses the real header, a firewall port/proto bypass. The fix makes parseV6 land
// on the same offset the host does. // on the same offset the host does.
func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) { func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
p := &firewall.Packet{} p := &firewall.ParsedPacket{}
const ( const (
hdrLen = 40 // IPv6 header 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. // 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") 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)
}
+8 -25
View File
@@ -1,28 +1,11 @@
package batch package batch
import "net/netip" // 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:
type RxBatcher interface { // a re-handshake replaces the tunnel outright and the replacement's epoch is higher,
// Reserve creates a pkt to borrow // so the old tunnel's packets sort first during the cutover overlap.
Reserve(sz int) []byte // Counter is the packet's AEAD message counter within that tunnel.
// Commit borrows pkt. The caller must keep pkt valid until the next Flush type SortKey struct {
Commit(pkt []byte) error Epoch uint64
// Flush emits every queued packet in arrival order. Counter uint64
// Returns the first error observed; keeps draining so one bad packet doesn't hold up the rest.
// After Flush returns, borrowed payload slices may be recycled.
Flush() error
}
type TxBatcher interface {
// Reserve creates a pkt to borrow
Reserve(sz int) []byte
// Commit borrows pkt and records its destination plus the 2-bit
// IP-level ECN codepoint to set on the outer (carrier) header. The
// caller must keep pkt valid until the next Flush. Pass 0 (Not-ECT)
// to leave the outer ECN field unset.
Commit(pkt []byte, dst netip.AddrPort, outerECN byte)
// Flush emits every queued packet via the underlying batch writer in arrival order.
// Returns an errors.Join of one or more errors.
// After Flush returns, borrowed payload slices may be recycled.
Flush() error
} }
+187
View File
@@ -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)
}
}
}
+96 -116
View File
@@ -8,135 +8,125 @@ import (
// flowKey identifies a transport flow by {src, dst, sport, dport, family}. // flowKey identifies a transport flow by {src, dst, sport, dport, family}.
// Comparable, so map lookups and linear scans over the slot list stay tight. // 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 // 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 // openSlots map, so a TCP and UDP flow on the same 5-tuple-without-proto never alias.
// never alias.
type flowKey struct { type flowKey struct {
src, dst [16]byte src, dst [16]byte
sport, dport uint16 sport, dport uint16
isV6 bool isV6 bool
} }
// initialSlots is the starting capacity of the slot pool. One flow per // initialSlots is the starting capacity of the slot pool.
// packet is the worst case so this matches a typical carrier-side // One flow per packet is the worst case,
// recvmmsg batch on the encrypted UDP socket. // so this matches a typical carrier-side recvmmsg batch on the UDP socket.
const initialSlots = 64 const initialSlots = 64
// parsedIP is the IP-level result of parseIPPrologue. The caller layers // parseIPAt validates the IP header for lane parsing. newPacket already resolved the L4 protocol
// L4-specific parsing (TCP / UDP) on top. // and offset for the firewall, so there is no proto sniff here; the caller's ipHdrLen is
type parsedIP struct { // cross-checked instead. A plain header (v4 IHL 20, v6 exactly 40) is the only coalesceable
fk flowKey // shape. The v6 check is load-bearing: it rejects extension-header packets whose L4 is not at
ipHdrLen int // byte 40.
// pkt is the original buffer trimmed to the IP-declared total length. //
// Anything below the IP layer (transport parsers) should slice into // The prologues fill fk's addresses and family in place (ports belong to the L4 parser; fk must
// pkt rather than the unbounded original. // be zero on entry so the v4 path leaves src[4:]/dst[4:] clear for map equality) and return pkt
pkt []byte // 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
} }
// parseIPPrologue extracts the IP-level fields the coalescers care about: // parseIPv4Prologue is the shared IPv4 tail of the prologue entries; the caller has verified
// IHL/payload length, version, src/dst addresses, and the L4 protocol byte. // len(pkt) >= 20 and the version.
// Returns ok=false for malformed input, IPv4 with options or fragmentation, func (fk *flowKey) parseIPv4Prologue(pkt []byte) ([]byte, bool) {
// or IPv6 with extension headers (all rejected by both coalescers in
// identical ways before this refactor).
//
// On success, p.pkt is len-trimmed to the IP-declared length so callers
// don't have to repeat the trim. wantProto is the IANA protocol number to
// require (6 for TCP, 17 for UDP); ok=false for any other value.
func parseIPPrologue(pkt []byte, wantProto byte) (parsedIP, bool) {
var p parsedIP
if len(pkt) < 20 {
return p, false
}
v := pkt[0] >> 4
switch v {
case 4:
ihl := int(pkt[0]&0x0f) * 4 ihl := int(pkt[0]&0x0f) * 4
if ihl != 20 { if ihl != 20 {
return p, false return nil, false
} }
if pkt[9] != wantProto { // Reject any fragmentation (MF or nonzero offset). The dispatcher already gated FragAny; kept
return p, false // as defense in depth, since a fragment folded into a superpacket would corrupt reassembly.
}
// Reject actual fragmentation (MF or non-zero frag offset).
if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 { if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 {
return p, false return nil, false
} }
totalLen := int(binary.BigEndian.Uint16(pkt[2:4])) totalLen := int(binary.BigEndian.Uint16(pkt[2:4]))
if totalLen > len(pkt) || totalLen < ihl { if totalLen > len(pkt) || totalLen < ihl {
return p, false return nil, false
}
p.ipHdrLen = 20
p.fk.isV6 = false
copy(p.fk.src[:4], pkt[12:16])
copy(p.fk.dst[:4], pkt[16:20])
p.pkt = pkt[:totalLen]
case 6:
if len(pkt) < 40 {
return p, false
}
if pkt[6] != wantProto {
return p, false
} }
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])) payloadLen := int(binary.BigEndian.Uint16(pkt[4:6]))
if 40+payloadLen > len(pkt) { if 40+payloadLen > len(pkt) {
return p, false return nil, false
} }
p.ipHdrLen = 40 fk.isV6 = true
p.fk.isV6 = true copy(fk.src[:], pkt[8:24])
copy(p.fk.src[:], pkt[8:24]) copy(fk.dst[:], pkt[24:40])
copy(p.fk.dst[:], pkt[24:40]) return pkt[:40+payloadLen], true
p.pkt = pkt[:40+payloadLen]
default:
return p, false
}
return p, true
} }
// ipHeadersMatch compares the IP portion of two packet header prefixes for // ipHeadersMatch compares the IP portion of two packet header prefixes for
// byte-for-byte equality on every field that must be identical across // byte-for-byte equality on every field that must be identical across coalesced segments.
// coalesced segments. Size/IPID/IPCsum are masked out. The full DSCP/ECN // Size/IPID/IPCsum are masked out.
// byte (IPv4 ToS / IPv6 traffic class) is compared, matching Linux kernel // The full DSCP/ECN byte (IPv4 ToS / IPv6 traffic class) is compared, matching Linux kernel GRO:
// GRO: segments with differing ECN codepoints must not coalesce, otherwise // segments with differing ECN codepoints must not coalesce,
// ORing e.g. ECT(0) with ECT(1) would fabricate a false CE (congestion) // 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.
// mark or mark a Not-ECT flow as ECN-capable.
// //
// The transport (L4) portion of the header is checked separately by the // The transport (L4) portion of the header is checked separately by the per-protocol matcher.
// per-protocol matcher.
func ipHeadersMatch(a, b []byte, isV6 bool) bool { func ipHeadersMatch(a, b []byte, isV6 bool) bool {
if isV6 { if isV6 {
// IPv6: byte 0 = version/TC[7:4], byte 1 = TC[3:0]/flow[19:16], // IPv6: [0:4] = version/TC/flow label (TC[1:0] is ECN, so the full TC byte must match),
// bytes [2:4] = flow[15:0], [6:8] = next_hdr/hop, [8:40] = src+dst. // [6:40] = next_hdr/hop + src + dst. Skip [4:6] payload_len.
// Compare byte 1 fully so ECN (TC[1:0]) must match. Skip [4:6] payload_len. return bytes.Equal(a[:4], b[:4]) && bytes.Equal(a[6:40], b[6:40])
if a[0] != b[0] {
return false
} }
if a[1] != b[1] { // IPv4: [0:2] = version/IHL + DSCP|ECN (full ECN byte must match),
return false // [6:10] = flags/fragoff/TTL/proto, [12:20] = src+dst.
}
if !bytes.Equal(a[2:4], b[2:4]) {
return false
}
if !bytes.Equal(a[6:40], b[6:40]) {
return false
}
return true
}
// IPv4: byte 0 = version/IHL, byte 1 = DSCP(6)|ECN(2),
// [6:10] flags/fragoff/TTL/proto, [12:20] src+dst.
// Compare byte 1 fully so ECN must match.
// Skip [2:4] total len, [4:6] id, [10:12] csum. // Skip [2:4] total len, [4:6] id, [10:12] csum.
if a[0] != b[0] { return bytes.Equal(a[:2], b[:2]) && bytes.Equal(a[6:10], b[6:10]) && bytes.Equal(a[12:20], b[12:20])
return false }
}
if a[1] != b[1] { // ipv4FlagDF is the Don't Fragment bit in the IPv4 flags byte (header byte 6).
return false const ipv4FlagDF = 0x40
}
if !bytes.Equal(a[6:10], b[6:10]) { // ipv4CanCoalesceID reports whether an IPv4 packet whose header starts at
return false // 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
if !bytes.Equal(a[12:20], b[12:20]) { // seed_id+n, so coalescing is only transparent when that re-stamp is either
return false // 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 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 // Arena is an injectable byte-slab that hands out non-overlapping borrowed
@@ -145,17 +135,14 @@ type Arena struct {
buf []byte buf []byte
} }
// NewArena returns an Arena with a pre-allocated backing of the given // NewArena returns an Arena with a pre-allocated backing of the given capacity.
// capacity. Pass 0 if you don't intend to call Reserve (e.g. a test that
// only feeds the coalescer pre-made []byte packets via Commit).
func NewArena(capacity int) *Arena { func NewArena(capacity int) *Arena {
return &Arena{buf: make([]byte, 0, capacity)} return &Arena{buf: make([]byte, 0, capacity)}
} }
// Reserve hands out a non-overlapping sz-byte slice from the arena. If the // Reserve hands out a non-overlapping sz-byte slice from the arena.
// request doesn't fit the current backing, a fresh, larger backing is // If the request doesn't fit the current backing, a fresh, larger backing is allocated.
// allocated; already-borrowed slices reference the old backing and remain // Already-borrowed slices reference the old backing and remain valid until Reset.
// valid until Reset.
func (a *Arena) Reserve(sz int) []byte { func (a *Arena) Reserve(sz int) []byte {
if len(a.buf)+sz > cap(a.buf) { if len(a.buf)+sz > cap(a.buf) {
newCap := max(cap(a.buf)*2, sz) newCap := max(cap(a.buf)*2, sz)
@@ -166,16 +153,9 @@ func (a *Arena) Reserve(sz int) []byte {
return a.buf[start : start+sz : start+sz] return a.buf[start : start+sz : start+sz]
} }
// Reset releases every slice handed out since the last Reset. Callers must // Reset releases every slice handed out since the last Reset.
// not use any previously-borrowed slice after this returns. The underlying // Callers must not use any previously-borrowed slice after this returns.
// backing array is retained so subsequent Reserves don't re-allocate. // The underlying backing array is retained so subsequent Reserves don't re-allocate.
func (a *Arena) Reset() { func (a *Arena) Reset() {
a.buf = a.buf[:0] a.buf = a.buf[:0]
} }
// Reserver hands out an sz-byte slice valid until its Resetter runs.
type Reserver func(sz int) []byte
// Resetter clears all reservations held by a Reserver. Only the arena's
// owner holds one; lanes inside a MultiCoalescer get nil.
type Resetter func()
+112
View File
@@ -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))
}
+76
View File
@@ -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)
}
+82 -81
View File
@@ -1,119 +1,121 @@
package batch package batch
import ( import (
"cmp"
"errors" "errors"
"io" "io"
"log/slog" "log/slog"
"slices"
"github.com/slackhq/nebula/firewall"
) )
// MultiCoalescer fans plaintext packets out to lane-specific batchers based // MultiCoalescer stages plaintext packets with their (epoch, counter) sort keys and, at Flush,
// on the IP/L4 protocol of the packet, sharing a single Reserve arena // replays them in sender-transmission order into lane-specific batchers selected by L4 protocol.
// across lanes so the caller's allocation pattern is unchanged.
// //
// Lanes are processed independently: the TCP coalescer only sees TCP, the // Sorting before dispatch keeps the ordering story simple: each lane consumes packets in
// UDP coalescer only sees UDP, and the passthrough lane handles everything // transmission order, builds slots in that order, and emits them in creation order. Wire reorder
// else. Per-flow arrival order is preserved because a single 5-tuple only // inside a flush batch is repaired here, before it can fragment a lane's coalesce chains, so the
// ever lands in one lane and each lane preserves its own slot order. // lanes carry no reorder-repair machinery.
// //
// Cross-lane order is NOT preserved across the TCP/UDP/passthrough split. // The contract is per-tunnel transmission order within each lane, with two exceptions: a pure TCP
// This is acceptable because the carrier-side recvmmsg path already // ACK may be overtaken by later same-flow data (it does not close the flow's open chain; a late
// stable-sorts by (peer, message counter) before delivering plaintext // ACK is just a stale ACK), and an unparseable shape seals every open chain in its lane (its flow
// here, so replay-window invariants are unaffected, and apps observe // is unknown) and rides the lane as an in-lane verbatim, still in transmission order. Routing
// correct per-flow ordering — which is all the IP layer guarantees anyway. // follows the flow: a flow's non-coalesceable shapes ride its protocol lane rather than falling
// Do not "fix" this by interleaving lane outputs at flush time; that // to the later-flushed pt lane.
// negates the entire point of coalescing (each lane needs to see runs of //
// adjacent same-flow packets to coalesce them). // Cross-lane order (TCP vs UDP vs everything else) is not preserved.
type MultiCoalescer struct { type MultiCoalescer struct {
tcp *TCPCoalescer tcp *TCPCoalescer
udp *UDPCoalescer udp *UDPCoalescer
pt *Passthrough pt *Passthrough
// arena is owned by the Multi: lanes get only its Reserve (nil Resetter)
// and Flush resets it exactly once after every lane has drained. // staged holds this batch's packets and sort keys until Flush. Borrowed: the caller keeps
arena *Arena // each pkt alive until Flush returns.
staged []stagedPacket
} }
// DefaultMultiArenaCap is the recommended arena capacity for a Multi-lane // stagedPacket carries the scalars dispatch needs from the firewall's ParsedPacket, copied by
// batcher: 64 slots × 65535 bytes ≈ 4 MiB, enough to hold one recvmmsg // value: pp is reused by the caller per packet and must not be retained past Commit.
// burst worth of MTU-sized packets without the arena growing. type stagedPacket struct {
const DefaultMultiArenaCap = initialSlots * 65535 pkt []byte
key SortKey
proto byte
fragAny bool
ipHdrLen uint16
}
// NewMultiCoalescer builds a multi-lane batcher. tcpEnabled lets the caller // NewMultiCoalescer builds a multi-lane batcher over w, based on available protocol support. The
// opt out of TCP coalescing (e.g. when the queue can't do TSO); udpEnabled // staging sort applies even when no GSO lane is available: passthrough-only platforms still get
// likewise gates UDP coalescing (only enable when USO was negotiated). // transmission-order repair.
// Either lane disabled redirects its traffic into the passthrough lane. func NewMultiCoalescer(w io.Writer, l *slog.Logger) *MultiCoalescer {
// arena is the single backing slab shared across every lane; the caller
// pre-sizes it via NewArena so the hot path never allocates.
func NewMultiCoalescer(w io.Writer, l *slog.Logger, arena *Arena, tcpEnabled, udpEnabled bool) *MultiCoalescer {
m := &MultiCoalescer{ m := &MultiCoalescer{
pt: NewPassthrough(w, arena.Reserve, nil), pt: NewPassthrough(w),
arena: arena, staged: make([]stagedPacket, 0, initialSlots),
}
if tcpEnabled {
m.tcp = NewTCPCoalescer(w, l, arena.Reserve, nil)
}
if udpEnabled {
m.udp = NewUDPCoalescer(w, arena.Reserve, nil)
} }
m.tcp = NewTCPCoalescer(w, l)
m.udp = NewUDPCoalescer(w)
return m return m
} }
func (m *MultiCoalescer) Reserve(sz int) []byte { // Commit stages pkt for the next Flush; dispatch is deferred so it runs on packets already in
return m.arena.Reserve(sz) // 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
} }
// Commit dispatches pkt to the appropriate lane based on IP version + L4 // compareStaged orders staged packets by (epoch, counter)
// proto. Borrowed slice contract is identical to the single-lane batchers, func compareStaged(a, b stagedPacket) int {
// pkt must remain valid until the next Flush. if c := cmp.Compare(a.key.Epoch, b.key.Epoch); c != 0 {
// return c
// On the success path the IP/TCP-or-UDP parse happens here once and the
// parsed struct is handed to the lane via commitParsed so the lane doesn't
// re-walk the header.
func (m *MultiCoalescer) Commit(pkt []byte) error {
if len(pkt) < 20 {
return m.pt.Commit(pkt)
} }
v := pkt[0] >> 4 return cmp.Compare(a.key.Counter, b.key.Counter)
var proto byte }
switch v {
case 4: // dispatch routes one staged packet to its protocol lane (see commitStaged), or to the verbatim
proto = pkt[9] // passthrough when the lane has no GSO support.
case 6: func (m *MultiCoalescer) dispatch(sp stagedPacket) error {
if len(pkt) < 40 { switch sp.proto {
return m.pt.Commit(pkt)
}
proto = pkt[6]
default:
return m.pt.Commit(pkt)
}
switch proto {
case ipProtoTCP: case ipProtoTCP:
if m.tcp != nil { if m.tcp != nil {
info, ok := parseTCPBase(pkt) return m.tcp.commitStaged(sp)
if !ok {
// Malformed/unsupported TCP shape (IP options, fragments, ...).
// Handle this via passthrough support in the TCP coalescer, to attempt to preserve flow order.
m.tcp.addPassthrough(pkt)
return nil
}
return m.tcp.commitParsed(pkt, info)
} }
case ipProtoUDP: case ipProtoUDP:
if m.udp != nil { if m.udp != nil {
info, ok := parseUDP(pkt) return m.udp.commitStaged(sp)
if !ok {
m.udp.addPassthrough(pkt) //we could also m.pt.Commit() here I guess?
return nil
}
return m.udp.commitParsed(pkt, info)
} }
} }
return m.pt.Commit(pkt) return m.pt.enqueue(sp.pkt)
} }
// Flush drains every lane in a fixed order, then resets the shared arena once. // Flush sorts the staged batch into transmission order, replays it into the lanes, then flushes each lane.
// A lane error doesn't stop the remaining lanes; the joined errors are returned. // 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 { 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 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 m.tcp != nil {
if err := m.tcp.Flush(); err != nil { if err := m.tcp.Flush(); err != nil {
errs = append(errs, err) errs = append(errs, err)
@@ -127,6 +129,5 @@ func (m *MultiCoalescer) Flush() error {
if err := m.pt.Flush(); err != nil { if err := m.pt.Flush(); err != nil {
errs = append(errs, err) errs = append(errs, err)
} }
m.arena.Reset()
return errors.Join(errs...) return errors.Join(errs...)
} }
+362 -21
View File
@@ -1,17 +1,39 @@
package batch package batch
import ( import (
"bytes"
"encoding/binary"
"io"
"testing" "testing"
"github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/test" "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 // TestMultiCoalescerRoutesByProto confirms TCP/UDP/other land in the right
// lane: TCP and UDP get coalesced when their lanes are enabled, anything // lane: TCP and UDP get coalesced when their lanes are enabled, anything
// else (ICMP here) falls through to plain Write. // else (ICMP here) falls through to plain Write.
func TestMultiCoalescerRoutesByProto(t *testing.T) { func TestMultiCoalescerRoutesByProto(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, true) m := newTestMultiCoalescer(t, w)
k := &keySeq{epoch: 1}
tcpPay := make([]byte, 1200) tcpPay := make([]byte, 1200)
udpPay := make([]byte, 1200) udpPay := make([]byte, 1200)
@@ -21,19 +43,19 @@ func TestMultiCoalescerRoutesByProto(t *testing.T) {
icmp[3] = 28 icmp[3] = 28
icmp[9] = 1 icmp[9] = 1
if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay)); err != nil { if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(1000, tcpAck, tcpPay))); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if err := m.Commit(buildTCPv4(2200, tcpAck, tcpPay)); err != nil { if err := m.Commit(buildTCPv4(2200, tcpAck, tcpPay), k.next(), testPP(buildTCPv4(2200, tcpAck, tcpPay))); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if err := m.Commit(buildUDPv4(2000, 53, udpPay)); err != nil { if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if err := m.Commit(buildUDPv4(2000, 53, udpPay)); err != nil { if err := m.Commit(buildUDPv4(2000, 53, udpPay), k.next(), testPP(buildUDPv4(2000, 53, udpPay))); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if err := m.Commit(icmp); err != nil { if err := m.Commit(icmp, k.next(), testPP(icmp)); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if err := m.Flush(); err != nil { if err := m.Flush(); err != nil {
@@ -48,17 +70,162 @@ func TestMultiCoalescerRoutesByProto(t *testing.T) {
} }
} }
// TestMultiCoalescerDisabledUDPFallsThrough verifies that when the UDP lane // TestMultiCoalescerRestoresTransmissionOrder is the core staging-sort
// is disabled (e.g. kernel doesn't support USO), UDP packets still reach // property: packets committed out of counter order (wire reorder inside one
// the kernel via the passthrough lane rather than being lost. // flush batch) are replayed into the lanes in transmission order, so the
func TestMultiCoalescerDisabledUDPFallsThrough(t *testing.T) { // 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} w := &fakeTunWriter{gsoEnabled: true}
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, false) // TSO on, USO off m := newTestMultiCoalescer(t, w)
pay := make([]byte, 1200)
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil { // 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) t.Fatal(err)
} }
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); 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}, 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) t.Fatal(err)
} }
if err := m.Flush(); err != nil { if err := m.Flush(); err != nil {
@@ -72,16 +239,164 @@ func TestMultiCoalescerDisabledUDPFallsThrough(t *testing.T) {
} }
} }
// TestMultiCoalescerDisabledTCPFallsThrough mirrors the TSO=off case. // TestMultiCoalescerNoOffloadsStillSorts covers a queue that can't offload
func TestMultiCoalescerDisabledTCPFallsThrough(t *testing.T) { // anything. Both lane constructors refuse, so every packet rides the
w := &fakeTunWriter{gsoEnabled: true} // verbatim lane — but the staging sort still applies, so emission follows
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), false, true) // TSO off, USO on // transmission order even without GSO.
func TestMultiCoalescerNoOffloadsStillSorts(t *testing.T) {
pay := make([]byte, 1200) w := &fakeTunWriter{gsoEnabled: false}
if err := m.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { 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) t.Fatal(err)
} }
if err := m.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil { }
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) t.Fatal(err)
} }
if err := m.Flush(); err != nil { if err := m.Flush(); err != nil {
@@ -94,3 +409,29 @@ func TestMultiCoalescerDisabledTCPFallsThrough(t *testing.T) {
t.Errorf("TCP must pass through as 2 plain writes, got %d", len(w.writes)) 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
}
+6 -31
View File
@@ -2,54 +2,29 @@ package batch
import ( import (
"io" "io"
"github.com/slackhq/nebula/udp"
) )
// Passthrough is a RxBatcher that doesn't batch anything, it just accumulates and then sends packets. // Passthrough is MultiCoalescer's verbatim lane: no batching, packets are written at Flush in the
// order enqueued.
type Passthrough struct { type Passthrough struct {
out io.Writer out io.Writer
slots [][]byte slots [][]byte
reserver Reserver
resetter Resetter
cursor int
} }
const passthroughBaseNumSlots = 128 func NewPassthrough(w io.Writer) *Passthrough {
// DefaultPassthroughArenaCap is the recommended arena capacity for a
// standalone Passthrough batcher: 128 slots × udp.MTU ≈ 1.1 MiB.
const DefaultPassthroughArenaCap = passthroughBaseNumSlots * udp.MTU
func NewPassthrough(w io.Writer, reserver Reserver, resetter Resetter) *Passthrough {
return &Passthrough{ return &Passthrough{
out: w, out: w,
slots: make([][]byte, 0, passthroughBaseNumSlots), slots: make([][]byte, 0, 128),
reserver: reserver,
resetter: resetter,
} }
} }
func (p *Passthrough) Reserve(sz int) []byte { // enqueue accepts one packet, already sorted into transmission order by dispatch.
return p.reserver(sz) func (p *Passthrough) enqueue(pkt []byte) error {
}
func (p *Passthrough) Commit(pkt []byte) error {
p.slots = append(p.slots, pkt) p.slots = append(p.slots, pkt)
return nil return nil
} }
// Flush drains every queued packet and calls the configured Resetter
func (p *Passthrough) Flush() error { func (p *Passthrough) Flush() error {
firstErr := p.drain()
if p.resetter != nil {
p.resetter()
}
return firstErr
}
// drain writes out every queued packet and clears the slot list.
func (p *Passthrough) drain() error {
var firstErr error var firstErr error
for _, s := range p.slots { for _, s := range p.slots {
_, err := p.out.Write(s) _, err := p.out.Write(s)
+182 -438
View File
@@ -2,12 +2,9 @@ package batch
import ( import (
"bytes" "bytes"
"context"
"encoding/binary" "encoding/binary"
"io" "io"
"log/slog" "log/slog"
"net/netip"
"slices"
"github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/overlay/tio"
) )
@@ -23,24 +20,19 @@ const tcpCoalesceBufSize = 65535
// superpacket. Keeping this well below the kernel's TSO ceiling bounds latency. // superpacket. Keeping this well below the kernel's TSO ceiling bounds latency.
const tcpCoalesceMaxSegs = 64 const tcpCoalesceMaxSegs = 64
// tcpCoalesceHdrCap is the scratch space we copy a seed's IP+TCP header // coalesceSlot is one entry in the coalescer's ordered event queue. A verbatim slot holds a single
// into. IPv6 (40) + TCP with full options (60) = 100 bytes. // borrowed packet emitted as-is (pure ACK, non-admissible TCP, unparseable, or oversize seed); a
const tcpCoalesceHdrCap = 100 // 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.
// coalesceSlot is one entry in the coalescer's ordered event queue. When
// passthrough is true the slot holds a single borrowed packet that must be
// emitted verbatim (non-TCP, non-admissible TCP, or oversize seed). When
// passthrough is false the slot is an in-progress coalesced superpacket:
// hdrBuf is a mutable copy of the seed's IP+TCP header (we patch total
// length and pseudo-header partial at flush), and payIovs are *borrowed*
// slices from the caller's plaintext buffers — no payload is ever copied.
// The caller (listenOut) must keep those buffers alive until Flush.
type coalesceSlot struct { type coalesceSlot struct {
passthrough bool verbatim bool
rawPkt []byte // borrowed when passthrough // 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 fk flowKey
hdrBuf [tcpCoalesceHdrCap]byte
hdrLen int hdrLen int
ipHdrLen int ipHdrLen int
isV6 bool isV6 bool
@@ -48,56 +40,49 @@ type coalesceSlot struct {
numSeg int numSeg int
totalPay int totalPay int
nextSeq uint32 nextSeq uint32
// psh closes the chain: set when the last-accepted segment had PSH or
// was sub-gsoSize. No further appends after that.
psh bool
payIovs [][]byte payIovs [][]byte
} }
// TCPCoalescer accumulates adjacent in-flow TCP data segments across // TCPCoalescer accumulates adjacent in-flow TCP data segments across multiple concurrent flows and
// multiple concurrent flows and emits each flow's run as a single TSO // emits each flow's run as a single TSO superpacket via tio.GSOWriter. Input must be in sender
// superpacket via tio.GSOWriter. All output — coalesced or not — is // transmission order (MultiCoalescer sorts by (epoch, counter) before dispatch); slots are emitted
// deferred until Flush so arrival order is preserved on the wire. Owns // in creation order, so emission reproduces transmission order except for the pure-ACK case in
// no locks; one coalescer per TUN write queue. // commitParsed. Owns no locks; one coalescer per TUN write queue.
type TCPCoalescer struct { type TCPCoalescer struct {
plainW io.Writer w tio.GSOWriter
gsoW tio.GSOWriter // nil when the queue doesn't support TSO
// slots is the ordered event queue. Flush walks it once and emits each // slots is the ordered event queue. Flush walks it once and emits each
// entry as either a WriteGSO (coalesced) or a plainW.Write (passthrough). // entry as either a WriteGSO (coalesced) or a w.Write (verbatim).
slots []*coalesceSlot slots []*coalesceSlot
// openSlots maps a flow key to its most recent non-sealed slot, so new // openSlots maps a flow key to its open slot so new segments can extend an in-progress
// segments can extend an in-progress superpacket in O(1). Slots are // superpacket in O(1). Removal is what closes a chain: on PSH or a short last segment, on a
// removed from this map when they close (PSH or short-last-segment), // non-admissible packet for the flow, or in Flush.
// when a non-admissible packet for that flow arrives, or in Flush.
openSlots map[flowKey]*coalesceSlot openSlots map[flowKey]*coalesceSlot
// lastSlot caches the most recently touched open slot. Steady-state // lastSlot caches the most recently touched open slot. Bulk traffic
// bulk traffic is dominated by a single flow, so comparing the // arrives in same-flow runs (single-flow steady state, or GRO bursts
// incoming key against the cached slot's own fk lets the hot path // under multi-flow), so comparing the incoming key against the cached
// skip the map lookup (and the aeshash of a 38-byte key) entirely. // 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 // Kept in lockstep with openSlots: nil whenever the slot it pointed
// at is removed/sealed. // at is removed.
lastSlot *coalesceSlot lastSlot *coalesceSlot
pool []*coalesceSlot // free list for reuse pool []*coalesceSlot // free list for reuse
reserver Reserver
resetter Resetter
l *slog.Logger l *slog.Logger
} }
func NewTCPCoalescer(w io.Writer, l *slog.Logger, reserver Reserver, resetter Resetter) *TCPCoalescer { // NewTCPCoalescer wraps w, returning nil if w can't accept GSO_TCP writes.
c := &TCPCoalescer{ func NewTCPCoalescer(w io.Writer, l *slog.Logger) *TCPCoalescer {
plainW: w, gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP)
if !ok {
return nil
}
return &TCPCoalescer{
w: gw,
slots: make([]*coalesceSlot, 0, initialSlots), slots: make([]*coalesceSlot, 0, initialSlots),
openSlots: make(map[flowKey]*coalesceSlot, initialSlots), openSlots: make(map[flowKey]*coalesceSlot, initialSlots),
pool: make([]*coalesceSlot, 0, initialSlots), pool: make([]*coalesceSlot, 0, initialSlots),
reserver: reserver,
resetter: resetter,
l: l, l: l,
} }
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP); ok {
c.gsoW = gw
}
return c
} }
// parsedTCP holds the fields extracted from a single parse so later steps // parsedTCP holds the fields extracted from a single parse so later steps
@@ -105,159 +90,153 @@ func NewTCPCoalescer(w io.Writer, l *slog.Logger, reserver Reserver, resetter Re
type parsedTCP struct { type parsedTCP struct {
fk flowKey fk flowKey
ipHdrLen int ipHdrLen int
tcpHdrLen int
hdrLen int hdrLen int
payLen int payLen int
seq uint32 seq uint32
flags byte flags byte
} }
// parseTCPBase extracts the flow key and IP/TCP offsets for any TCP packet, // parseAt extracts the flow key and IP/TCP offsets for a packet the dispatcher already knows is
// regardless of whether it's admissible for coalescing. Returns ok=false // TCP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on
// for non-TCP or malformed input. // entry and is filled in place; see flowKey.parseIPAt for why. Returns false for malformed input
// Accepts IPv4 (no options or fragmentation) and IPv6 (no extension headers). // or any shape that must not coalesce (IPv4 options/fragmentation, IPv6 extension headers).
func parseTCPBase(pkt []byte) (parsedTCP, bool) { func (p *parsedTCP) parseAt(pkt []byte, ipHdrLen int) bool {
var p parsedTCP trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen)
ip, ok := parseIPPrologue(pkt, ipProtoTCP)
if !ok { if !ok {
return p, false return false
} }
pkt = ip.pkt return p.parseTail(trimmed, ipHdrLen)
p.fk = ip.fk
p.ipHdrLen = ip.ipHdrLen
if len(pkt) < p.ipHdrLen+20 {
return p, false
}
tcpOff := int(pkt[p.ipHdrLen+12]>>4) * 4
if tcpOff < 20 || tcpOff > 60 {
return p, false
}
if len(pkt) < p.ipHdrLen+tcpOff {
return p, false
}
p.tcpHdrLen = tcpOff
p.hdrLen = p.ipHdrLen + tcpOff
p.payLen = len(pkt) - p.hdrLen
p.seq = binary.BigEndian.Uint32(pkt[p.ipHdrLen+4 : p.ipHdrLen+8])
p.flags = pkt[p.ipHdrLen+13]
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2])
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
return p, true
} }
// TCP flag bits (byte 13 of the TCP header). Only the bits actually consulted // parseTail layers the TCP-header parse on a validated IP prologue. pkt is the trimmed packet;
// by the coalescer are named; FIN/SYN/RST/URG/CWR are rejected via the // fk's addresses are already filled.
// negative mask in coalesceable, not by name. 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 ( const (
tcpFlagPsh = 0x08 tcpFlagPsh = 0x08
tcpFlagAck = 0x10 tcpFlagAck = 0x10
tcpFlagEce = 0x40 tcpFlagEce = 0x40
) )
// coalesceable reports whether a parsed TCP segment is eligible for // sealAllOpen closes every open coalesce chain. Called for unparseable packets: the flow key is
// coalescing. Accepts ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE with a // unknown, so any open chain could otherwise absorb later data and emit it ahead of this packet.
// non-empty payload. CWR is excluded because it marks a one-shot func (c *TCPCoalescer) sealAllOpen() {
// congestion-window-reduced transition the receiver must observe at a clear(c.openSlots)
// segment boundary. c.lastSlot = nil
func (p parsedTCP) coalesceable() bool {
if p.flags&tcpFlagAck == 0 {
return false
}
if p.flags&^(tcpFlagAck|tcpFlagPsh|tcpFlagEce) != 0 {
return false
}
return p.payLen > 0
} }
func (c *TCPCoalescer) Reserve(sz int) []byte { // sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips
return c.reserver(sz) // 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 {
// Commit borrows pkt. The caller must keep pkt valid until the next Flush. return
func (c *TCPCoalescer) Commit(pkt []byte) error {
if c.gsoW == nil {
c.addPassthrough(pkt)
return nil
} }
info, ok := parseTCPBase(pkt) if last := c.lastSlot; last != nil && last.fk == fk {
if !ok {
c.addPassthrough(pkt)
return nil
}
return c.commitParsed(pkt, info)
}
// commitParsed is the post-parse half of Commit. The caller must have
// already verified parseTCPBase succeeded (info is a valid TCP parse).
// Used by MultiCoalescer.Commit to avoid re-walking the IP/TCP header
// after the dispatcher has already done so.
func (c *TCPCoalescer) commitParsed(pkt []byte, info parsedTCP) error {
if c.gsoW == nil {
c.addPassthrough(pkt)
return nil
}
if !info.coalesceable() {
// TCP but not admissible (SYN/FIN/RST/URG/CWR or zero-payload).
// Seal this flow's open slot so later in-flow packets don't extend
// it and accidentally reorder past this passthrough.
if last := c.lastSlot; last != nil && last.fk == info.fk {
c.lastSlot = nil c.lastSlot = nil
} }
delete(c.openSlots, info.fk) delete(c.openSlots, fk)
c.addPassthrough(pkt) }
// 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 return nil
} }
// Single-flow fast path: with only one open flow the cache hits every // Cached-slot fast path. Arrival isn't per-packet interleaved even with
// packet, and len(openSlots)==1 lets us skip the 38-byte fk compare // many flows: wire-side GRO delivers runs of same-flow packets
// when there are multiple flows in flight (where the hit rate would // (deliverSegments splits a superdatagram into up to 64), so the cache
// be ~0 and the compare is pure overhead). // hits for the length of each run and a miss costs one fk compare
// before the map lookup carries the weight.
var open *coalesceSlot var open *coalesceSlot
if last := c.lastSlot; last != nil && len(c.openSlots) == 1 && last.fk == info.fk { if last := c.lastSlot; last != nil && last.fk == info.fk {
open = last open = last
} else { } else {
open = c.openSlots[info.fk] open = c.openSlots[info.fk]
} }
if open != nil { if open != nil {
if c.canAppend(open, pkt, info) { if c.canAppend(open, pkt, info) {
c.appendPayload(open, pkt, info) if c.appendPayload(open, pkt, info) {
if open.psh { // Chain closed (PSH or short segment): stop extending it.
delete(c.openSlots, info.fk) c.sealFlow(info.fk)
c.lastSlot = nil
} else { } else {
c.lastSlot = open c.lastSlot = open
} }
return nil return nil
} }
// Can't extend — seal it and fall through to seed a fresh slot. // Can't extend (seq gap from upstream loss, header change, or a full
delete(c.openSlots, info.fk) // chain): evict it from openSlots and fall through to seed a fresh slot.
if c.lastSlot == open { c.sealFlow(info.fk)
c.lastSlot = nil
}
} }
c.seed(pkt, info) c.seed(pkt, info)
return nil return nil
} }
// Flush emits every queued event in (per-flow) seq order.
func (c *TCPCoalescer) Flush() error { func (c *TCPCoalescer) Flush() error {
first := c.drain()
if c.resetter != nil {
c.resetter()
}
return first
}
// drain emits every queued slot (reordering/merging coalesced runs first)
// and clears the slot state.
func (c *TCPCoalescer) drain() error {
c.reorderForFlush()
var first error var first error
for _, s := range c.slots { for _, s := range c.slots {
var err error var err error
if s.passthrough { if s.verbatim || s.numSeg == 1 {
_, err = c.plainW.Write(s.rawPkt) // 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 { } else {
err = c.flushSlot(s) err = c.flushSlot(s)
} }
@@ -274,23 +253,27 @@ func (c *TCPCoalescer) drain() error {
return first return first
} }
func (c *TCPCoalescer) addPassthrough(pkt []byte) { func (c *TCPCoalescer) addVerbatim(pkt []byte) {
s := c.take() s := c.take()
s.passthrough = true s.verbatim = true
s.rawPkt = pkt s.rawPkt = pkt
c.slots = append(c.slots, s) c.slots = append(c.slots, s)
} }
func (c *TCPCoalescer) seed(pkt []byte, info parsedTCP) { func (c *TCPCoalescer) seed(pkt []byte, info *parsedTCP) {
if info.hdrLen > tcpCoalesceHdrCap || info.hdrLen+info.payLen > tcpCoalesceBufSize { if info.hdrLen+info.payLen > tcpCoalesceBufSize {
// Pathological shape can't fit our scratch, emit as-is. // Pathological shape that can't ride a superpacket; emit as-is. No chain for this flow can
c.addPassthrough(pkt) // 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 return
} }
s := c.take() s := c.take()
s.passthrough = false s.verbatim = false
s.rawPkt = nil // rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is
copy(s.hdrBuf[:], pkt[:info.hdrLen]) // the superpacket header flushSlot patches in place.
s.rawPkt = pkt
s.hdrLen = info.hdrLen s.hdrLen = info.hdrLen
s.ipHdrLen = info.ipHdrLen s.ipHdrLen = info.ipHdrLen
s.isV6 = info.fk.isV6 s.isV6 = info.fk.isV6
@@ -299,26 +282,23 @@ func (c *TCPCoalescer) seed(pkt []byte, info parsedTCP) {
s.numSeg = 1 s.numSeg = 1
s.totalPay = info.payLen s.totalPay = info.payLen
s.nextSeq = info.seq + uint32(info.payLen) s.nextSeq = info.seq + uint32(info.payLen)
s.psh = info.flags&tcpFlagPsh != 0
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen]) s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
c.slots = append(c.slots, s) c.slots = append(c.slots, s)
if !s.psh { if info.flags&tcpFlagPsh == 0 {
c.openSlots[info.fk] = s c.openSlots[info.fk] = s
c.lastSlot = s c.lastSlot = s
} else if last := c.lastSlot; last != nil && last.fk == info.fk { } else {
// PSH-on-seed seals the slot immediately. Any prior cached open // PSH on the seed closes the chain immediately; it is never registered as open.
// slot for this flow has just been sealed-and-replaced by this // Drop any stale entry for this flow too (defense in depth, unreachable if lastSlot's lockstep invariant holds).
// passthrough-shaped seed, so drop the cache too. c.sealFlow(info.fk)
c.lastSlot = nil
} }
} }
// canAppend reports whether info's packet extends the slot's seed: same // canAppend reports whether info's packet extends the slot's seed: same header shape and stable
// header shape and stable contents, adjacent seq, not oversized, chain not closed. // contents, adjacent seq, not oversized. A closed chain never reaches here; closing removes the
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info parsedTCP) bool { // slot from openSlots, the only path in. The header fields read from rawPkt are always pristine:
if s.psh { // the only pre-flush mutation is the PSH propagate, which also closes the chain.
return false func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info *parsedTCP) bool {
}
if info.hdrLen != s.hdrLen { if info.hdrLen != s.hdrLen {
return false return false
} }
@@ -334,31 +314,35 @@ func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info parsedTCP) bo
if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize { if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize {
return false return false
} }
// ECE state must be stable across a burst — receivers expect the // ECE state must be stable across a burst.
// flag set on every segment of a CE-echoing window or none. // Receivers expect the flag set on every segment of a CE-echoing window or none.
seedFlags := s.hdrBuf[s.ipHdrLen+13] seedFlags := s.rawPkt[s.ipHdrLen+13]
if (seedFlags^info.flags)&tcpFlagEce != 0 { if (seedFlags^info.flags)&tcpFlagEce != 0 {
return false return false
} }
if !headersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) { 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 false
} }
return true return true
} }
func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info parsedTCP) { // 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.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
s.numSeg++ s.numSeg++
s.totalPay += info.payLen s.totalPay += info.payLen
s.nextSeq = info.seq + uint32(info.payLen) s.nextSeq = info.seq + uint32(info.payLen)
if info.flags&tcpFlagPsh != 0 { if info.flags&tcpFlagPsh != 0 {
// Propagate PSH into the seed header so kernel TSO sets it on the // Propagate PSH into the seed header so kernel TSO sets it on the last segment. Mutating
// last segment. Without this the sender's push signal is dropped. // rawPkt is safe: PSH also closes the chain, so no admission check re-reads this header.
s.hdrBuf[s.ipHdrLen+13] |= tcpFlagPsh s.rawPkt[s.ipHdrLen+13] |= tcpFlagPsh
}
if info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0 {
s.psh = true
} }
return info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0
} }
func (c *TCPCoalescer) take() *coalesceSlot { func (c *TCPCoalescer) take() *coalesceSlot {
@@ -372,21 +356,18 @@ func (c *TCPCoalescer) take() *coalesceSlot {
} }
func (c *TCPCoalescer) release(s *coalesceSlot) { func (c *TCPCoalescer) release(s *coalesceSlot) {
s.passthrough = false
s.rawPkt = nil
clear(s.payIovs) clear(s.payIovs)
s.payIovs = s.payIovs[:0] *s = coalesceSlot{payIovs: s.payIovs[:0]}
s.numSeg = 0
s.totalPay = 0
s.psh = false
c.pool = append(c.pool, s) c.pool = append(c.pool, s)
} }
// flushSlot patches the header and calls WriteGSO. Does not remove the slot from c.slots. // 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 { func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
total := s.hdrLen + s.totalPay total := s.hdrLen + s.totalPay
l4Len := total - s.ipHdrLen l4Len := total - s.ipHdrLen
hdr := s.hdrBuf[:s.hdrLen] hdr := s.rawPkt[:s.hdrLen]
if s.isV6 { if s.isV6 {
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len)) binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
@@ -406,7 +387,7 @@ func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
tcsum := s.ipHdrLen + 16 tcsum := s.ipHdrLen + 16
binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum)) binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum))
return c.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoTCP) 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 // headersMatch compares two IP+TCP header prefixes for byte-for-byte
@@ -437,242 +418,6 @@ func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
return true return true
} }
// reorderForFlush neutralizes wire-side reorder that the rxOrder buffer
// couldn't catch (anything crossing a recvmmsg batch boundary). Without
// this pass a small wire reorder — counter 250 arriving in batch K when
// 200..249 are coming in batch K+1 — would seed an out-of-seq slot first
// and emit it ahead of the lower-seq slot, manifesting at the inner TCP
// receiver as a much larger reorder than the wire actually had.
//
// Two phases:
// 1. Sort each passthrough-bounded segment of c.slots by (flow, seq).
// Cross-flow ordering inside a segment isn't preserved (it never was
// and doesn't matter for any single flow's TCP correctness).
// 2. Sweep once and merge adjacent same-flow slots whose ranges are now
// contiguous AND whose tail is gsoSize-aligned. The tail constraint
// matters because the kernel TSO splitter chops at gsoSize from the
// start of the merged payload — a short segment in the middle would
// desynchronize every later segment.
//
// Passthrough slots act as barriers: the merge check skips them on either
// side, so a SYN/FIN/RST/CWR is never reordered relative to its flow's
// data.
func (c *TCPCoalescer) reorderForFlush() {
if len(c.slots) <= 1 {
return
}
runStart := 0
for i := 0; i <= len(c.slots); i++ {
if i < len(c.slots) && !c.slots[i].passthrough {
continue
}
c.sortRun(c.slots[runStart:i])
runStart = i + 1
}
out := c.slots[:0]
logged := false
for _, s := range c.slots {
if n := len(out); n > 0 {
prev := out[n-1]
if !prev.passthrough && !s.passthrough && prev.fk == s.fk {
// Same-flow neighbors after sort. If they aren't seq-
// contiguous it's a real gap — packets the wire reordered
// across batches, or actual loss before nebula. Log it so
// the operator can quantify how often it happens; the data
// itself still emits in seq order, kernel TCP handles the
// gap via its OOO queue.
if c.l.Enabled(context.Background(), slog.LevelDebug) {
if prev.nextSeq != slotSeedSeq(s) {
logged = true
gap := int64(slotSeedSeq(s)) - int64(prev.nextSeq)
c.l.Debug("tcp coalesce: cross-slot seq gap",
"src", flowKeyAddr(s.fk, false),
"dst", flowKeyAddr(s.fk, true),
"sport", s.fk.sport,
"dport", s.fk.dport,
"prev_seed_seq", slotSeedSeq(prev),
"prev_next_seq", prev.nextSeq,
"this_seed_seq", slotSeedSeq(s),
"gap_bytes", gap,
"prev_seg_count", prev.numSeg,
"prev_total_pay", prev.totalPay,
)
}
}
if canMergeSlots(prev, s) {
mergeSlots(prev, s)
c.release(s)
continue
}
}
}
out = append(out, s)
}
if logged {
c.l.Warn("==== end of batch ====")
}
c.slots = out
}
// flowKeyAddr returns the src or dst address from fk as a netip.Addr for
// logging. Only used on the cold gap-log path so the netip allocation
// doesn't matter.
func flowKeyAddr(fk flowKey, dst bool) netip.Addr {
src := fk.src
if dst {
src = fk.dst
}
if fk.isV6 {
return netip.AddrFrom16(src)
}
var v4 [4]byte
copy(v4[:], src[:4])
return netip.AddrFrom4(v4)
}
// sortRun stable-sorts run by (flowKey, seedSeq) so each flow's slots
// cluster together in seq order, ready for the merge sweep. Stable so
// equal-key slots keep their original relative position (defensive — a
// duplicate seedSeq would already mean something's wrong upstream).
func (c *TCPCoalescer) sortRun(run []*coalesceSlot) {
if len(run) <= 1 {
return
}
// slices.SortStableFunc with a free, non-capturing comparator avoids the
// reflection + closure-escape allocations that sort.SliceStable forces.
slices.SortStableFunc(run, compareCoalesceSlots)
}
func compareCoalesceSlots(a, b *coalesceSlot) int {
if cmp := flowKeyCompare(a.fk, b.fk); cmp != 0 {
return cmp
}
aSeq, bSeq := slotSeedSeq(a), slotSeedSeq(b)
if aSeq == bSeq {
return 0
}
if tcpSeqLess(aSeq, bSeq) {
return -1
}
return 1
}
// slotSeedSeq returns the TCP seq of the slot's seed (first segment).
// nextSeq tracks the seq just past the last appended byte; subtracting
// totalPay walks back to the seed. uint32 wraparound is the right TCP
// arithmetic so no special-casing is needed.
func slotSeedSeq(s *coalesceSlot) uint32 {
return s.nextSeq - uint32(s.totalPay)
}
// tcpSeqLess reports whether a precedes b in TCP serial-number arithmetic
// (RFC 1323 §2.3). The signed int32 cast turns the modular subtraction
// into the right comparison even across the 2^32 wrap.
func tcpSeqLess(a, b uint32) bool {
return int32(a-b) < 0
}
// flowKeyCompare orders flowKeys deterministically. The exact ordering
// is irrelevant — only that same-flow slots cluster together so the
// post-sort sweep can merge contiguous pairs.
func flowKeyCompare(a, b flowKey) int {
// Cheap scalar fields first so most non-matching keys short-circuit
// without ever calling bytes.Compare. sport is the ephemeral port on
// egress flows and discriminates fastest. For matching keys (same
// flow), array equality on src/dst inlines to word-sized compares,
// so we only pay bytes.Compare when the arrays actually differ.
if a.sport != b.sport {
if a.sport < b.sport {
return -1
}
return 1
}
if a.dport != b.dport {
if a.dport < b.dport {
return -1
}
return 1
}
if a.dst != b.dst {
return bytes.Compare(a.dst[:], b.dst[:])
}
if a.src != b.src {
return bytes.Compare(a.src[:], b.src[:])
}
if a.isV6 != b.isV6 {
if !a.isV6 {
return -1
}
return 1
}
return 0
}
// canMergeSlots reports whether s can fold into prev as one merged TSO
// superpacket. Same flow, contiguous TCP byte range, equal gsoSize, and
// fits within the kernel TSO limits. The tail-of-prev check rejects any
// merge whose first slot ended on a sub-gsoSize segment — kernel TSO
// would split the merged skb at gsoSize boundaries from the start, so a
// short segment in the middle would corrupt every later segment. PSH and
// ECE state must agree across both slots: PSH is a semantic delimiter
// (preserving the sender's push boundary) and ECE state must be uniform
// across a window (the same rule canAppend enforces for in-flow appends).
// The IP-level ECN codepoint must also match: this check calls headersMatch
// → ipHeadersMatch, which compares the full DSCP/ECN byte, so two slots with
// differing ECN marks stay separate superpackets, each keeping its own mark.
//
// Note: a slot sealed by reorder (canAppend returned false on seq
// mismatch) keeps psh=false, so this restriction does not block the
// reorder-fix merge — only legitimate PSH-set seals.
func canMergeSlots(prev, s *coalesceSlot) bool {
if prev.psh {
return false
}
if prev.fk != s.fk {
return false
}
if prev.gsoSize != s.gsoSize {
return false
}
if prev.nextSeq != slotSeedSeq(s) {
return false
}
if prev.numSeg+s.numSeg > tcpCoalesceMaxSegs {
return false
}
if prev.hdrLen+prev.totalPay+s.totalPay > tcpCoalesceBufSize {
return false
}
if len(prev.payIovs[len(prev.payIovs)-1]) != prev.gsoSize {
return false
}
prevFlags := prev.hdrBuf[prev.ipHdrLen+13]
sFlags := s.hdrBuf[s.ipHdrLen+13]
if (prevFlags^sFlags)&tcpFlagEce != 0 {
return false
}
if !headersMatch(prev.hdrBuf[:prev.hdrLen], s.hdrBuf[:s.hdrLen], prev.isV6, prev.ipHdrLen) {
return false
}
return true
}
// mergeSlots folds src into dst in place: payIovs concatenated, counters
// and totals updated, PSH OR'd into the seed header so the push signal is
// not lost. The seed header's seq, gsoSize, and fk are unchanged. Caller
// is responsible for releasing src (it's no longer in c.slots after this call).
func mergeSlots(dst, src *coalesceSlot) {
dst.payIovs = append(dst.payIovs, src.payIovs...)
dst.numSeg += src.numSeg
dst.totalPay += src.totalPay
dst.nextSeq = src.nextSeq
if src.psh {
dst.psh = true
dst.hdrBuf[dst.ipHdrLen+13] |= tcpFlagPsh
}
}
// ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must // ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must
// already have its checksum field zeroed) and returns the folded/inverted // already have its checksum field zeroed) and returns the folded/inverted
// 16-bit value to store. // 16-bit value to store.
@@ -717,9 +462,8 @@ func pseudoSumIPv6(src, dst []byte, proto byte, l4Len int) uint32 {
return sum return sum
} }
// foldOnceNoInvert folds the 32-bit accumulator to 16 bits and returns it // foldOnceNoInvert folds the 32-bit accumulator to 16 bits and returns it unchanged (no one's complement).
// unchanged (no one's complement). This is what virtio NEEDS_CSUM wants in // This is what virtio NEEDS_CSUM wants in the L4 checksum field
// the L4 checksum field — the kernel will add the payload sum and invert.
func foldOnceNoInvert(sum uint32) uint16 { func foldOnceNoInvert(sum uint32) uint16 {
for sum>>16 != 0 { for sum>>16 != 0 {
sum = (sum & 0xffff) + (sum >> 16) sum = (sum & 0xffff) + (sum >> 16)
+51 -78
View File
@@ -2,9 +2,9 @@ package batch
import ( import (
"encoding/binary" "encoding/binary"
"runtime"
"testing" "testing"
"github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/overlay/tio"
"github.com/slackhq/nebula/test" "github.com/slackhq/nebula/test"
) )
@@ -55,7 +55,31 @@ func buildTCPv4Interleaved(nFlows, perFlow, payloadLen int) [][]byte {
return pkts return pkts
} }
// buildICMPv4 returns a minimal non-TCP packet that takes the passthrough // 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. // branch in Commit.
func buildICMPv4() []byte { func buildICMPv4() []byte {
pkt := make([]byte, 28) pkt := make([]byte, 28)
@@ -71,8 +95,7 @@ func buildICMPv4() []byte {
// between batches, and reports per-packet cost. // between batches, and reports per-packet cost.
func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) { func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
b.Helper() b.Helper()
arena := NewArena(0) c := newTestTCPCoalescer(b, nopTunWriter{})
c := NewTCPCoalescer(nopTunWriter{}, test.NewLogger(), arena.Reserve, arena.Reset)
b.ReportAllocs() b.ReportAllocs()
b.SetBytes(int64(len(pkts[0]))) b.SetBytes(int64(len(pkts[0])))
b.ResetTimer() b.ResetTimer()
@@ -113,8 +136,17 @@ func BenchmarkCommitInterleaved16(b *testing.B) {
runCommitBench(b, pkts, len(pkts)) runCommitBench(b, pkts, len(pkts))
} }
// BenchmarkCommitPassthrough exercises the non-TCP branch: parseTCPBase // BenchmarkCommitRunInterleaved4 is 4 concurrent flows arriving in
// bails early and addPassthrough is the only work. // 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) { func BenchmarkCommitPassthrough(b *testing.B) {
pkt := buildICMPv4() pkt := buildICMPv4()
pkts := make([][]byte, 64) pkts := make([][]byte, 64)
@@ -126,7 +158,7 @@ func BenchmarkCommitPassthrough(b *testing.B) {
// BenchmarkCommitNonCoalesceableTCP sends SYN|ACK packets on one flow. // BenchmarkCommitNonCoalesceableTCP sends SYN|ACK packets on one flow.
// Each packet takes the "TCP but not admissible" branch which does a // Each packet takes the "TCP but not admissible" branch which does a
// map delete + passthrough. Measures the seal-without-slot cost. // map delete + verbatim. Measures the seal-without-slot cost.
func BenchmarkCommitNonCoalesceableTCP(b *testing.B) { func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
pay := make([]byte, 0) pay := make([]byte, 0)
pkts := make([][]byte, 64) pkts := make([][]byte, 64)
@@ -136,18 +168,24 @@ func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
runCommitBench(b, pkts, 64) runCommitBench(b, pkts, 64)
} }
// runMultiCommitBench drives MultiCoalescer.Commit. The dispatcher does // runMultiCommitBench drives MultiCoalescer.Commit with in-order keys, so
// the IP/L4 parse once and passes the parsed struct to the lane, so this // it includes the staging sort's already-sorted fast path plus the
// is the bench that shows the savings of skipping the lane's re-parse. // 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) { func runMultiCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
b.Helper() b.Helper()
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger(), NewArena(0), true, true) m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger())
pps := make([]*firewall.ParsedPacket, len(pkts))
for i, p := range pkts {
pps[i] = testPP(p)
}
b.ReportAllocs() b.ReportAllocs()
b.SetBytes(int64(len(pkts[0]))) b.SetBytes(int64(len(pkts[0])))
b.ResetTimer() b.ResetTimer()
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
pkt := pkts[i%len(pkts)] j := i % len(pkts)
if err := m.Commit(pkt); err != nil { if err := m.Commit(pkts[j], SortKey{Epoch: 1, Counter: uint64(i + 1)}, pps[j]); err != nil {
b.Fatal(err) b.Fatal(err)
} }
if (i+1)%batchSize == 0 { if (i+1)%batchSize == 0 {
@@ -174,68 +212,3 @@ func BenchmarkMultiCommitInterleaved4(b *testing.B) {
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200) pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
runMultiCommitBench(b, pkts, len(pkts)) runMultiCommitBench(b, pkts, len(pkts))
} }
// flowKeyPair is one comparison input for the flowKeyCompare bench.
type flowKeyPair struct{ a, b flowKey }
// makeFlowKey builds an IPv4 flowKey from compact inputs.
func makeFlowKey(srcLow, dstLow uint32, sport, dport uint16) flowKey {
var fk flowKey
binary.BigEndian.PutUint32(fk.src[12:16], srcLow)
binary.BigEndian.PutUint32(fk.dst[12:16], dstLow)
fk.sport = sport
fk.dport = dport
return fk
}
// flowKeyCases are the workload mixes flowKeyCompare sees in practice.
// - sameFlow: equal keys; tests the equal-path cost (sort runs hit this
// repeatedly when many segments share a flow).
// - sportDiffers: same src/dst/dport, different sport — the typical
// "sibling flows from one host to one server" pattern.
// - dstDiffers: same src/sport/dport, different dst — outbound to many
// servers from a fixed local port.
// - allDiffer: every field differs; worst case for short-circuiting.
func flowKeyCases() map[string][]flowKeyPair {
const n = 64
cases := map[string][]flowKeyPair{
"sameFlow": make([]flowKeyPair, n),
"sportDiffers": make([]flowKeyPair, n),
"dstDiffers": make([]flowKeyPair, n),
"allDiffer": make([]flowKeyPair, n),
}
for i := range n {
base := makeFlowKey(0x0a000001, 0x0a000002, 40000, 443)
cases["sameFlow"][i] = flowKeyPair{a: base, b: base}
cases["sportDiffers"][i] = flowKeyPair{
a: base,
b: makeFlowKey(0x0a000001, 0x0a000002, uint16(40001+i), 443),
}
cases["dstDiffers"][i] = flowKeyPair{
a: base,
b: makeFlowKey(0x0a000001, uint32(0x0a000002+i+1), 40000, 443),
}
cases["allDiffer"][i] = flowKeyPair{
a: makeFlowKey(uint32(0x0a000001+i), uint32(0x0a000002+i), uint16(40000+i), uint16(80+i)),
b: makeFlowKey(uint32(0x0b000001+i), uint32(0x0b000002+i), uint16(50000+i), uint16(443+i)),
}
}
return cases
}
// BenchmarkFlowKeyCompare measures flowKeyCompare across the workloads
// the sort step actually sees. Use this to compare reorderings.
func BenchmarkFlowKeyCompare(b *testing.B) {
for name, pairs := range flowKeyCases() {
b.Run(name, func(b *testing.B) {
b.ReportAllocs()
b.ResetTimer()
var sink int
for i := 0; i < b.N; i++ {
p := pairs[i&(len(pairs)-1)]
sink += flowKeyCompare(p.a, p.b)
}
runtime.KeepAlive(sink)
})
}
}
File diff suppressed because it is too large Load Diff
+8 -9
View File
@@ -6,7 +6,7 @@ const SendBatchCap = 128
// batchWriter is the minimal subset of udp.Conn needed by SendBatch to flush. // batchWriter is the minimal subset of udp.Conn needed by SendBatch to flush.
type batchWriter interface { type batchWriter interface {
WriteBatch(bufs [][]byte, addrs []netip.AddrPort, outerECNs []byte) error WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error)
} }
// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch. // SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch.
@@ -16,7 +16,6 @@ type SendBatch struct {
out batchWriter out batchWriter
bufs [][]byte bufs [][]byte
dsts []netip.AddrPort dsts []netip.AddrPort
ecns []byte
arena *Arena arena *Arena
} }
@@ -26,7 +25,6 @@ func NewSendBatch(out batchWriter, batchCap, arenaSize int) *SendBatch {
out: out, out: out,
bufs: make([][]byte, 0, batchCap), bufs: make([][]byte, 0, batchCap),
dsts: make([]netip.AddrPort, 0, batchCap), dsts: make([]netip.AddrPort, 0, batchCap),
ecns: make([]byte, 0, batchCap),
arena: NewArena(arenaSize), arena: NewArena(arenaSize),
} }
} }
@@ -40,21 +38,22 @@ func (b *SendBatch) Reserve(sz int) []byte {
// bounding how long the first packet of a large read batch waits. // bounding how long the first packet of a large read batch waits.
func (b *SendBatch) Len() int { return len(b.bufs) } func (b *SendBatch) Len() int { return len(b.bufs) }
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort, outerECN byte) { func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort) {
b.bufs = append(b.bufs, pkt) b.bufs = append(b.bufs, pkt)
b.dsts = append(b.dsts, dst) b.dsts = append(b.dsts, dst)
b.ecns = append(b.ecns, outerECN)
} }
func (b *SendBatch) Flush() error { // 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 var err error
written := 0
if len(b.bufs) > 0 { if len(b.bufs) > 0 {
err = b.out.WriteBatch(b.bufs, b.dsts, b.ecns) written, err = b.out.WriteBatch(b.bufs, b.dsts)
} }
clear(b.bufs) clear(b.bufs)
b.bufs = b.bufs[:0] b.bufs = b.bufs[:0]
b.dsts = b.dsts[:0] b.dsts = b.dsts[:0]
b.ecns = b.ecns[:0]
b.arena.Reset() b.arena.Reset()
return err return written, err
} }
+10 -12
View File
@@ -8,10 +8,9 @@ import (
type fakeBatchWriter struct { type fakeBatchWriter struct {
bufs [][]byte bufs [][]byte
addrs []netip.AddrPort addrs []netip.AddrPort
ecns []byte
} }
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, ecns []byte) error { func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
// Snapshot — SendBatch.Flush nils its slot pointers right after WriteBatch // Snapshot — SendBatch.Flush nils its slot pointers right after WriteBatch
// returns, so tests must capture data before that happens. // returns, so tests must capture data before that happens.
w.bufs = make([][]byte, len(bufs)) w.bufs = make([][]byte, len(bufs))
@@ -21,8 +20,7 @@ func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, ecns
w.bufs[i] = cp w.bufs[i] = cp
} }
w.addrs = append(w.addrs[:0], addrs...) w.addrs = append(w.addrs[:0], addrs...)
w.ecns = append(w.ecns[:0], ecns...) return len(bufs), nil
return nil
} }
func TestSendBatchReserveCommitFlush(t *testing.T) { func TestSendBatchReserveCommitFlush(t *testing.T) {
@@ -36,9 +34,9 @@ func TestSendBatchReserveCommitFlush(t *testing.T) {
t.Fatalf("slot %d: cap=%d want 32", i, cap(slot)) t.Fatalf("slot %d: cap=%d want 32", i, cap(slot))
} }
pkt := append(slot[:0], byte(i), byte(i+1), byte(i+2)) pkt := append(slot[:0], byte(i), byte(i+1), byte(i+2))
b.Commit(pkt, ap, 0) b.Commit(pkt, ap)
} }
if err := b.Flush(); err != nil { if _, err := b.Flush(); err != nil {
t.Fatalf("Flush: %v", err) t.Fatalf("Flush: %v", err)
} }
if len(fw.bufs) != 4 { if len(fw.bufs) != 4 {
@@ -55,7 +53,7 @@ func TestSendBatchReserveCommitFlush(t *testing.T) {
// Flush again with nothing committed — should be a no-op. // Flush again with nothing committed — should be a no-op.
fw.bufs = nil fw.bufs = nil
if err := b.Flush(); err != nil { if _, err := b.Flush(); err != nil {
t.Fatalf("empty Flush: %v", err) t.Fatalf("empty Flush: %v", err)
} }
if fw.bufs != nil { if fw.bufs != nil {
@@ -77,9 +75,9 @@ func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
for i := 0; i < 3; i++ { for i := 0; i < 3; i++ {
s := b.Reserve(8) s := b.Reserve(8)
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i)) pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
b.Commit(pkt, ap, 0) b.Commit(pkt, ap)
} }
if err := b.Flush(); err != nil { if _, err := b.Flush(); err != nil {
t.Fatalf("Flush: %v", err) t.Fatalf("Flush: %v", err)
} }
@@ -98,18 +96,18 @@ func TestSendBatchGrowPreservesCommitted(t *testing.T) {
s1 := b.Reserve(4) s1 := b.Reserve(4)
pkt1 := append(s1[:0], 0x11, 0x22, 0x33, 0x44) pkt1 := append(s1[:0], 0x11, 0x22, 0x33, 0x44)
b.Commit(pkt1, ap, 0) b.Commit(pkt1, ap)
s2 := b.Reserve(8) // exceeds remaining cap, triggers grow s2 := b.Reserve(8) // exceeds remaining cap, triggers grow
pkt2 := append(s2[:0], 0xA, 0xB, 0xC, 0xD, 0xE) pkt2 := append(s2[:0], 0xA, 0xB, 0xC, 0xD, 0xE)
b.Commit(pkt2, ap, 0) b.Commit(pkt2, ap)
// pkt1 must still be intact even though backing reallocated. // pkt1 must still be intact even though backing reallocated.
if pkt1[0] != 0x11 || pkt1[3] != 0x44 { if pkt1[0] != 0x11 || pkt1[3] != 0x44 {
t.Fatalf("first packet corrupted by grow: %x", pkt1) t.Fatalf("first packet corrupted by grow: %x", pkt1)
} }
if err := b.Flush(); err != nil { if _, err := b.Flush(); err != nil {
t.Fatalf("Flush: %v", err) t.Fatalf("Flush: %v", err)
} }
if len(fw.bufs) != 2 { if len(fw.bufs) != 2 {
+140 -150
View File
@@ -1,6 +1,7 @@
package batch package batch
import ( import (
"bytes"
"encoding/binary" "encoding/binary"
"io" "io"
@@ -18,72 +19,56 @@ const udpCoalesceBufSize = 65535
// accepts up to 64 segments per skb (UDP_MAX_SEGMENTS); stay under that. // accepts up to 64 segments per skb (UDP_MAX_SEGMENTS); stay under that.
const udpCoalesceMaxSegs = 64 const udpCoalesceMaxSegs = 64
// udpCoalesceHdrCap is the scratch space we copy a seed's IP+UDP header
// into. IPv6 (40) + UDP (8) = 48; round up for safety.
const udpCoalesceHdrCap = 64
// udpSlot is one entry in the UDPCoalescer's ordered event queue. // udpSlot is one entry in the UDPCoalescer's ordered event queue.
type udpSlot struct { type udpSlot struct {
passthrough bool verbatim bool
rawPkt []byte // borrowed when passthrough // 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 fk flowKey
hdrBuf [udpCoalesceHdrCap]byte
hdrLen int hdrLen int
ipHdrLen int ipHdrLen int
isV6 bool isV6 bool
gsoSize int // per-segment UDP payload length gsoSize int // per-segment UDP payload length
numSeg int numSeg int
totalPay int totalPay int
// sealed closes the chain: set when a sub-gsoSize segment is appended
// (kernel UDP-GSO requires every segment but the last to be exactly
// gsoSize) or when limits are hit. No further appends after.
sealed bool
payIovs [][]byte payIovs [][]byte
} }
// UDPCoalescer accumulates adjacent in-flow UDP datagrams across multiple // UDPCoalescer accumulates adjacent in-flow UDP datagrams across multiple
// concurrent flows and emits each flow's run as a single GSO_UDP_L4 // concurrent flows and emits each flow's run as a single GSO_UDP_L4 superpacket via tio.GSOWriter.
// superpacket via tio.GSOWriter. Falls back to per-packet writes when the // Preserves the in-flow order of packets as they are Commit-ed
// underlying writer doesn't support USO.
//
// All output — coalesced or not — is deferred until Flush so per-flow
// arrival order is preserved on the wire. Cross-flow order is NOT preserved
// across the TCP/UDP/passthrough split when this coalescer runs alongside
// others — see multi_coalesce.go. Per-flow order is preserved because a
// single 5-tuple only ever lands in one lane and each lane preserves its
// own slot order.
// //
// Owns no locks; one coalescer per TUN write queue. // Owns no locks; one coalescer per TUN write queue.
type UDPCoalescer struct { type UDPCoalescer struct {
plainW io.Writer w tio.GSOWriter
gsoW tio.GSOWriter // nil when the queue can't accept GSO_UDP_L4
slots []*udpSlot slots []*udpSlot
openSlots map[flowKey]*udpSlot openSlots map[flowKey]*udpSlot
// lastSlot caches the most recently touched open slot; see the
// TCPCoalescer field of the same name. Single-flow QUIC bulk is the
// dominant USO workload, and multi-flow arrival comes in GRO runs, so
// the fk compare beats the map's 38-byte key hash on most packets.
// Kept in lockstep with openSlots: nil whenever the slot it pointed at
// is removed.
lastSlot *udpSlot
pool []*udpSlot pool []*udpSlot
reserver Reserver
resetter Resetter
} }
// NewUDPCoalescer wraps w. The caller is responsible for only constructing func NewUDPCoalescer(w io.Writer) *UDPCoalescer {
// this when the underlying Queue's Capabilities advertise USO; otherwise gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP)
// the kernel may reject GSO_UDP_L4 writes. If w does not implement if !ok {
// tio.GSOWriter at all (single-packet Queue), the coalescer degrades to return nil
// plain Writes — same defensive shape as the TCP coalescer. }
func NewUDPCoalescer(w io.Writer, reserver Reserver, resetter Resetter) *UDPCoalescer { return &UDPCoalescer{
c := &UDPCoalescer{ w: gw,
plainW: w,
slots: make([]*udpSlot, 0, initialSlots), slots: make([]*udpSlot, 0, initialSlots),
openSlots: make(map[flowKey]*udpSlot, initialSlots), openSlots: make(map[flowKey]*udpSlot, initialSlots),
pool: make([]*udpSlot, 0, initialSlots), pool: make([]*udpSlot, 0, initialSlots),
reserver: reserver,
resetter: resetter,
} }
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP); ok {
c.gsoW = gw
}
return c
} }
// parsedUDP holds the fields extracted from a single parse so later steps // parsedUDP holds the fields extracted from a single parse so later steps
@@ -95,104 +80,111 @@ type parsedUDP struct {
payLen int payLen int
} }
// parseUDP extracts the flow key and IP/UDP offsets for a UDP packet. // parseAt extracts the flow key and IP/UDP offsets for a packet the dispatcher already knows is
// Returns ok=false for non-UDP, malformed, or unsupported header shapes // UDP; ipHdrLen is the upstream-resolved L4 offset (see flowKey.parseIPAt). p must be zero on
// (IPv4 with options/fragmentation, IPv6 with extension headers). // entry and is filled in place. Returns false for malformed input or any shape that must not
func parseUDP(pkt []byte) (parsedUDP, bool) { // coalesce (IPv4 options/fragmentation, IPv6 extension headers).
var p parsedUDP func (p *parsedUDP) parseAt(pkt []byte, ipHdrLen int) bool {
ip, ok := parseIPPrologue(pkt, ipProtoUDP) trimmed, ok := p.fk.parseIPAt(pkt, ipHdrLen)
if !ok { if !ok {
return p, false return false
} }
pkt = ip.pkt return p.parseTail(trimmed, ipHdrLen)
p.fk = ip.fk }
p.ipHdrLen = ip.ipHdrLen
if len(pkt) < p.ipHdrLen+8 { // parseTail layers the UDP-header parse on a validated IP prologue. pkt is the trimmed packet;
return p, false // fk's addresses are already filled.
func (p *parsedUDP) parseTail(pkt []byte, ipHdrLen int) bool {
if len(pkt) < ipHdrLen+8 {
return false
} }
p.hdrLen = p.ipHdrLen + 8
// UDP `length` field: must equal IP-derived length-of-UDP-header-plus-payload. // UDP `length` field: must equal IP-derived length-of-UDP-header-plus-payload.
udpLen := int(binary.BigEndian.Uint16(pkt[p.ipHdrLen+4 : p.ipHdrLen+6])) udpLen := int(binary.BigEndian.Uint16(pkt[ipHdrLen+4 : ipHdrLen+6]))
if udpLen < 8 || udpLen > len(pkt)-p.ipHdrLen { if udpLen < 8 || udpLen > len(pkt)-ipHdrLen {
return p, false return false
} }
p.ipHdrLen = ipHdrLen
p.hdrLen = ipHdrLen + 8
p.payLen = udpLen - 8 p.payLen = udpLen - 8
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2]) p.fk.sport = binary.BigEndian.Uint16(pkt[ipHdrLen : ipHdrLen+2])
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4]) p.fk.dport = binary.BigEndian.Uint16(pkt[ipHdrLen+2 : ipHdrLen+4])
return p, true return true
} }
func (c *UDPCoalescer) Reserve(sz int) []byte { // sealFlow closes fk's open chain, if any, keeping lastSlot in lockstep. The len guard skips
return c.reserver(sz) // 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)
} }
// Commit borrows pkt. The caller must keep pkt valid until the next Flush. // commitStaged commits one staged packet dispatch routed to this lane. A shape the lane cannot
func (c *UDPCoalescer) Commit(pkt []byte) error { // coalesce (any fragmentation, unparseable header) seals every open chain — its flow is unknown —
if c.gsoW == nil { // and rides the lane as an in-lane verbatim, still in transmission order.
c.addPassthrough(pkt) func (c *UDPCoalescer) commitStaged(sp stagedPacket) error {
if sp.fragAny {
c.sealAllOpen()
c.addVerbatim(sp.pkt)
return nil return nil
} }
info, ok := parseUDP(pkt) var info parsedUDP
if !ok { if !info.parseAt(sp.pkt, int(sp.ipHdrLen)) {
c.addPassthrough(pkt) c.sealAllOpen()
c.addVerbatim(sp.pkt)
return nil return nil
} }
return c.commitParsed(pkt, info) return c.commitParsed(sp.pkt, &info)
} }
// commitParsed is the post-parse half of Commit. The caller must have // commitParsed commits one parsed UDP packet. The caller (dispatch, via parseAt) supplies a
// already verified parseUDP succeeded. Used by MultiCoalescer.Commit to // valid parse so the header is not re-walked here.
// avoid re-walking the IP/UDP header. func (c *UDPCoalescer) commitParsed(pkt []byte, info *parsedUDP) error {
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
if c.gsoW == nil { // coalesced.
c.addPassthrough(pkt)
return nil
}
// A zero-length UDP datagram (UDP `length` == 8) is legal and must still
// reach the TUN, but it can't be coalesced: a GSO slot would store an
// empty payload iovec and the kernel has nothing to segment. Seal any
// open chain for this flow (so a later, non-empty datagram seeds fresh
// *after* this one and per-flow arrival order is preserved) and deliver
// it as a plain single datagram.
if info.payLen == 0 { if info.payLen == 0 {
delete(c.openSlots, info.fk) c.sealFlow(info.fk)
c.addPassthrough(pkt) c.addVerbatim(pkt)
return nil return nil
} }
if open := c.openSlots[info.fk]; open != 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.canAppend(open, pkt, info) {
c.appendPayload(open, pkt, info) if c.appendPayload(open, pkt, info) {
if open.sealed { // Chain closed (short segment): stop extending it.
delete(c.openSlots, info.fk) c.sealFlow(info.fk)
} else {
c.lastSlot = open
} }
return nil return nil
} }
// Can't extend — seal it and fall through to seed a fresh slot. // Can't extend: evict it from openSlots and fall through to seed a
delete(c.openSlots, info.fk) // fresh slot.
c.sealFlow(info.fk)
} }
c.seed(pkt, info) c.seed(pkt, info)
return nil return nil
} }
// Flush drains every queued slot and calls the configured Resetter.
func (c *UDPCoalescer) Flush() error { func (c *UDPCoalescer) Flush() error {
first := c.drain()
if c.resetter != nil {
c.resetter()
}
return first
}
// drain emits every queued slot in arrival order and clears the slot state.
// It does NOT reset the arena: borrowed payload slices stay valid until the
// arena's owner recycles it.
func (c *UDPCoalescer) drain() error {
var first error var first error
for _, s := range c.slots { for _, s := range c.slots {
var err error var err error
if s.passthrough { if s.verbatim || s.numSeg == 1 {
_, err = c.plainW.Write(s.rawPkt) // 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 { } else {
err = c.flushSlot(s) err = c.flushSlot(s)
} }
@@ -204,25 +196,38 @@ func (c *UDPCoalescer) drain() error {
clear(c.slots) clear(c.slots)
c.slots = c.slots[:0] c.slots = c.slots[:0]
clear(c.openSlots) clear(c.openSlots)
c.lastSlot = nil
return first return first
} }
func (c *UDPCoalescer) addPassthrough(pkt []byte) { // 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 := c.take()
s.passthrough = true s.verbatim = true
s.rawPkt = pkt s.rawPkt = pkt
c.slots = append(c.slots, s) c.slots = append(c.slots, s)
} }
func (c *UDPCoalescer) seed(pkt []byte, info parsedUDP) { func (c *UDPCoalescer) seed(pkt []byte, info *parsedUDP) {
if info.hdrLen > udpCoalesceHdrCap || info.hdrLen+info.payLen > udpCoalesceBufSize { if info.hdrLen+info.payLen > udpCoalesceBufSize {
c.addPassthrough(pkt) // 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 return
} }
s := c.take() s := c.take()
s.passthrough = false s.verbatim = false
s.rawPkt = nil // rawPkt serves the numSeg==1 fast path in Flush, is the header source for canAppend, and is
copy(s.hdrBuf[:], pkt[:info.hdrLen]) // the superpacket header flushSlot patches in place.
s.rawPkt = pkt
s.hdrLen = info.hdrLen s.hdrLen = info.hdrLen
s.ipHdrLen = info.ipHdrLen s.ipHdrLen = info.ipHdrLen
s.isV6 = info.fk.isV6 s.isV6 = info.fk.isV6
@@ -230,19 +235,16 @@ func (c *UDPCoalescer) seed(pkt []byte, info parsedUDP) {
s.gsoSize = info.payLen s.gsoSize = info.payLen
s.numSeg = 1 s.numSeg = 1
s.totalPay = info.payLen s.totalPay = info.payLen
s.sealed = false
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen]) s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
c.slots = append(c.slots, s) c.slots = append(c.slots, s)
c.openSlots[info.fk] = s c.openSlots[info.fk] = s
c.lastSlot = s
} }
// canAppend reports whether info's packet extends the slot's seed. // canAppend reports whether info's packet extends the slot's seed.
// Kernel UDP-GSO requires every segment except possibly the last to be // Kernel UDP-GSO requires every segment except possibly the last to be
// exactly gsoSize, and the last may be shorter (≤ gsoSize). // exactly gsoSize, and the last may be shorter (≤ gsoSize).
func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info parsedUDP) bool { func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info *parsedUDP) bool {
if s.sealed {
return false
}
if info.hdrLen != s.hdrLen { if info.hdrLen != s.hdrLen {
return false return false
} }
@@ -255,20 +257,25 @@ func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info parsedUDP) bool {
if s.hdrLen+s.totalPay+info.payLen > udpCoalesceBufSize { if s.hdrLen+s.totalPay+info.payLen > udpCoalesceBufSize {
return false return false
} }
if !udpHeadersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) { // 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 false
} }
return true return true
} }
func (c *UDPCoalescer) appendPayload(s *udpSlot, pkt []byte, info parsedUDP) { // 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.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
s.numSeg++ s.numSeg++
s.totalPay += info.payLen s.totalPay += info.payLen
if info.payLen < s.gsoSize { return info.payLen < s.gsoSize
// Last-segment-can-be-shorter: this seals the chain.
s.sealed = true
}
} }
func (c *UDPCoalescer) take() *udpSlot { func (c *UDPCoalescer) take() *udpSlot {
@@ -282,30 +289,19 @@ func (c *UDPCoalescer) take() *udpSlot {
} }
func (c *UDPCoalescer) release(s *udpSlot) { func (c *UDPCoalescer) release(s *udpSlot) {
s.passthrough = false // Reset every field, identity ones included; see TCPCoalescer.release.
s.rawPkt = nil
clear(s.payIovs) clear(s.payIovs)
s.payIovs = s.payIovs[:0] *s = udpSlot{payIovs: s.payIovs[:0]}
s.numSeg = 0
s.totalPay = 0
s.sealed = false
c.pool = append(c.pool, s) c.pool = append(c.pool, s)
} }
// flushSlot patches the IP header total length / IPv6 payload length and // 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 length to the *total* across all coalesced segments, then seeds
// the UDP checksum field with the pseudo-header partial (single-fold, not // the UDP checksum field with the pseudo-header partial (single-fold, not
// inverted) per virtio NEEDS_CSUM. The kernel's ip_rcv_core (v4) and // inverted) per virtio NEEDS_CSUM. The patches land in place in rawPkt; the
// ip6_rcv_core (v6) trim the skb to those length fields, so per-segment // slot is released right after, so nothing re-reads the patched header.
// values would silently drop everything but the first segment. The kernel
// then walks each segment in __udp_gso_segment, recomputing per-segment
// uh->len / iph->tot_len / IPv6 plen and adjusting the checksum via
// `check = csum16_add(csum16_sub(uh->check, uh->len), newlen)` — meaning
// our seed's uh->check must be consistent with the seed's uh->len, which
// is what passing the total to both pseudoSum and the UDP length field
// guarantees.
func (c *UDPCoalescer) flushSlot(s *udpSlot) error { func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
hdr := s.hdrBuf[:s.hdrLen] hdr := s.rawPkt[:s.hdrLen]
total := s.hdrLen + s.totalPay // full IP+UDP+all_payloads bytes total := s.hdrLen + s.totalPay // full IP+UDP+all_payloads bytes
l4Len := total - s.ipHdrLen // total UDP (8 + sum of payloads) l4Len := total - s.ipHdrLen // total UDP (8 + sum of payloads)
@@ -330,14 +326,11 @@ func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
udpCsumOff := s.ipHdrLen + 6 udpCsumOff := s.ipHdrLen + 6
binary.BigEndian.PutUint16(hdr[udpCsumOff:udpCsumOff+2], foldOnceNoInvert(psum)) binary.BigEndian.PutUint16(hdr[udpCsumOff:udpCsumOff+2], foldOnceNoInvert(psum))
return c.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoUDP) 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 // udpHeadersMatch compares two IP+UDP header prefixes for byte-equality on
// every field that must be identical across coalesced segments. Length // every field that must be identical across coalesced segments
// fields are masked out (flushSlot rewrites them), but the IP-level ECN
// codepoint is compared (via ipHeadersMatch) so segments with differing ECN
// don't coalesce, matching kernel GRO.
func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool { func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
if len(a) != len(b) { if len(a) != len(b) {
return false return false
@@ -345,11 +338,8 @@ func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
if !ipHeadersMatch(a, b, isV6) { if !ipHeadersMatch(a, b, isV6) {
return false return false
} }
// UDP: compare sport+dport ([0:4]). Skip length [4:6] and checksum [6:8] // UDP: compare sport+dport ([0:4]). Skip length [4:6] and checksum [6:8]:
// length varies (we rewrite at flush) and the checksum will be redone. // length varies (we rewrite at flush) and the checksum will be redone.
udp := ipHdrLen udp := ipHdrLen
if a[udp] != b[udp] || a[udp+1] != b[udp+1] || a[udp+2] != b[udp+2] || a[udp+3] != b[udp+3] { return bytes.Equal(a[udp:udp+4], b[udp:udp+4])
return false
}
return true
} }
+72
View File
@@ -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))
}
+134 -70
View File
@@ -1,7 +1,9 @@
package batch package batch
import ( import (
"bytes"
"encoding/binary" "encoding/binary"
"io"
"testing" "testing"
) )
@@ -58,29 +60,31 @@ func buildUDPv6(sport, dport uint16, payload []byte) []byte {
return pkt return pkt
} }
func TestUDPCoalescerPassthroughWhenGSOUnavailable(t *testing.T) { // newTestUDPCoalescer builds a coalescer over w and fails the test if w can't
w := &fakeTunWriter{gsoEnabled: false} // do USO. See newTestTCPCoalescer.
arena := NewArena(0) func newTestUDPCoalescer(tb testing.TB, w io.Writer) *UDPCoalescer {
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset) tb.Helper()
pkt := buildUDPv4(1000, 53, make([]byte, 100)) c := NewUDPCoalescer(w)
if err := c.Commit(pkt); err != nil { if c == nil {
t.Fatal(err) tb.Fatal("NewUDPCoalescer: writer does not support USO")
} }
if len(w.writes) != 0 || len(w.gsoWrites) != 0 { return c
t.Fatalf("no Add-time writes: writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) }
// 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 err := c.Flush(); err != nil { if c := NewUDPCoalescer(&plainOnlyWriter{}); c != nil {
t.Fatal(err) t.Fatalf("want nil for a plain writer, got %v", c)
}
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
t.Fatalf("want single plain write, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
} }
} }
func TestUDPCoalescerNonUDPPassthrough(t *testing.T) { func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := newTestUDPCoalescer(t, w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
// ICMP packet // ICMP packet
pkt := make([]byte, 28) pkt := make([]byte, 28)
pkt[0] = 0x45 pkt[0] = 0x45
@@ -101,8 +105,7 @@ func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) { func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := newTestUDPCoalescer(t, w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv4(1000, 53, make([]byte, 800)) pkt := buildUDPv4(1000, 53, make([]byte, 800))
if err := c.Commit(pkt); err != nil { if err := c.Commit(pkt); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -110,17 +113,21 @@ func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
if err := c.Flush(); err != nil { if err := c.Flush(); err != nil {
t.Fatal(err) t.Fatal(err)
} }
// Single-segment flush goes through WriteGSO; the writer infers GSO_NONE // A slot that never grew past one datagram flushes as a plain Write of
// from len(pays)==1 and the kernel fills in the UDP csum (NEEDS_CSUM). // the original packet bytes: the original (already valid) checksum
if len(w.gsoWrites) != 1 || len(w.writes) != 0 { // 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)) 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) { func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := newTestUDPCoalescer(t, w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
for i := 0; i < 3; i++ { for i := 0; i < 3; i++ {
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil { if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
@@ -160,8 +167,7 @@ func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
// Last segment may be shorter, sealing the chain. // Last segment may be shorter, sealing the chain.
func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) { func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := newTestUDPCoalescer(t, w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
full := make([]byte, 1200) full := make([]byte, 1200)
tail := make([]byte, 600) tail := make([]byte, 600)
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil { if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
@@ -180,22 +186,23 @@ func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
if err := c.Flush(); err != nil { if err := c.Flush(); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if len(w.gsoWrites) != 2 { // The sealed 3-datagram chain is a real superpacket; the re-seed stays
t.Fatalf("want 2 gso writes (sealed + new seed), got %d", len(w.gsoWrites)) // 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 { if len(w.gsoWrites[0].pays) != 3 {
t.Errorf("first super: want 3 pays, got %d", len(w.gsoWrites[0].pays)) t.Errorf("super: want 3 pays, got %d", len(w.gsoWrites[0].pays))
} }
if len(w.gsoWrites[1].pays) != 1 { if got, want := len(w.writes[0]), 20+8+1200; got != want {
t.Errorf("second super: want 1 pay (re-seed), got %d", len(w.gsoWrites[1].pays)) t.Errorf("re-seed plain write len=%d want %d", got, want)
} }
} }
// A larger-than-gsoSize packet cannot extend the slot — it reseeds. // A larger-than-gsoSize packet cannot extend the slot — it reseeds.
func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) { func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := newTestUDPCoalescer(t, w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil { if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -205,16 +212,21 @@ func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
if err := c.Flush(); err != nil { if err := c.Flush(); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if len(w.gsoWrites) != 2 { // Both seeds stay single-segment → two plain writes in arrival order.
t.Fatalf("want 2 separate seeds, got %d", len(w.gsoWrites)) 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. // Different 5-tuples must not coalesce.
func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) { func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := newTestUDPCoalescer(t, w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 800) pay := make([]byte, 800)
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil { if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -245,8 +257,7 @@ func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
// Caps at udpCoalesceMaxSegs. // Caps at udpCoalesceMaxSegs.
func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) { func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := newTestUDPCoalescer(t, w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 100) pay := make([]byte, 100)
for i := 0; i < udpCoalesceMaxSegs+5; i++ { for i := 0; i < udpCoalesceMaxSegs+5; i++ {
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil { if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
@@ -271,12 +282,12 @@ func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
// Differing IP ECN codepoints must not coalesce: udpHeadersMatch compares // Differing IP ECN codepoints must not coalesce: udpHeadersMatch compares
// the full ToS byte (matching kernel GRO). A CE-marked datagram mid-run // the full ToS byte (matching kernel GRO). A CE-marked datagram mid-run
// seals the Not-ECT chain and seeds a fresh superpacket that keeps CE; the // seals the Not-ECT chain and reseeds; the trailing Not-ECT datagram
// trailing Not-ECT datagram seeds another. // 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) { func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := newTestUDPCoalescer(t, w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 800) pay := make([]byte, 800)
pkt0 := buildUDPv4(1000, 53, pay) // ECN=00 (Not-ECT) pkt0 := buildUDPv4(1000, 53, pay) // ECN=00 (Not-ECT)
pkt1 := buildUDPv4(1000, 53, pay) pkt1 := buildUDPv4(1000, 53, pay)
@@ -290,16 +301,13 @@ func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
if err := c.Flush(); err != nil { if err := c.Flush(); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if len(w.gsoWrites) != 3 { if len(w.writes) != 3 || len(w.gsoWrites) != 0 {
t.Fatalf("want 3 separate seeds (differing ECN), got %d (plain=%d)", len(w.gsoWrites), len(w.writes)) 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} wantECN := []byte{0x00, 0x03, 0x00}
for i, g := range w.gsoWrites { for i, p := range w.writes {
if len(g.pays) != 1 { if got := p[1] & 0x03; got != wantECN[i] {
t.Errorf("gso %d pay count=%d want 1", i, len(g.pays)) t.Errorf("write %d ECN=%#x want %#x", i, got, wantECN[i])
}
if got := g.hdr[1] & 0x03; got != wantECN[i] {
t.Errorf("gso %d ECN=%#x want %#x", i, got, wantECN[i])
} }
} }
} }
@@ -307,8 +315,7 @@ func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
// IPv6 path: same flow, equal-sized → coalesced. // IPv6 path: same flow, equal-sized → coalesced.
func TestUDPCoalescerIPv6Coalesces(t *testing.T) { func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := newTestUDPCoalescer(t, w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
for i := 0; i < 3; i++ { for i := 0; i < 3; i++ {
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil { if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
@@ -344,8 +351,7 @@ func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
// DSCP differences must reseed: udpHeadersMatch compares the full ToS byte. // DSCP differences must reseed: udpHeadersMatch compares the full ToS byte.
func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) { func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := newTestUDPCoalescer(t, w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 800) pay := make([]byte, 800)
pkt0 := buildUDPv4(1000, 53, pay) pkt0 := buildUDPv4(1000, 53, pay)
pkt1 := buildUDPv4(1000, 53, pay) pkt1 := buildUDPv4(1000, 53, pay)
@@ -359,16 +365,16 @@ func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
if err := c.Flush(); err != nil { if err := c.Flush(); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if len(w.gsoWrites) != 2 { // Both seeds stay single-segment → two plain writes, no gso.
t.Fatalf("want 2 separate seeds (different DSCP), got %d", len(w.gsoWrites)) 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. // Fragmented IPv4 must not be coalesced.
func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) { func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := newTestUDPCoalescer(t, w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv4(1000, 53, make([]byte, 200)) pkt := buildUDPv4(1000, 53, make([]byte, 200))
binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1 binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1
if err := c.Commit(pkt); err != nil { if err := c.Commit(pkt); err != nil {
@@ -389,8 +395,7 @@ func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
// reach the GSO path. Regression: must not panic and must be written. // reach the GSO path. Regression: must not panic and must be written.
func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) { func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := newTestUDPCoalescer(t, w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv4(1000, 53, nil) // UDP length 8, zero payload pkt := buildUDPv4(1000, 53, nil) // UDP length 8, zero payload
if err := c.Commit(pkt); err != nil { if err := c.Commit(pkt); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -406,11 +411,10 @@ func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) {
} }
} }
// IPv6 zero-length UDP datagram: same passthrough contract as v4. // IPv6 zero-length UDP datagram: same verbatim contract as v4.
func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) { func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := newTestUDPCoalescer(t, w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv6(1000, 53, nil) // UDP length 8, zero payload pkt := buildUDPv6(1000, 53, nil) // UDP length 8, zero payload
if err := c.Commit(pkt); err != nil { if err := c.Commit(pkt); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -431,8 +435,7 @@ func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
// wire — per-flow arrival order (full, empty, full) must be preserved. // wire — per-flow arrival order (full, empty, full) must be preserved.
func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) { func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := newTestUDPCoalescer(t, w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
full := make([]byte, 800) full := make([]byte, 800)
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil { if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -447,17 +450,23 @@ func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
// The empty datagram sealed the first slot, so the trailing full packet // The empty datagram sealed the first slot, so the trailing full packet
// can't join it: two single-segment superpackets bracket one plain write. // can't join it. All three emit as plain writes (the two full datagrams
if len(w.gsoWrites) != 2 || len(w.writes) != 1 { // stayed single-segment; the empty one is verbatim) in per-flow
t.Fatalf("want 2 gso writes + 1 plain, got gso=%d plain=%d", len(w.gsoWrites), len(w.writes)) // 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). // IPv4 with options is not admissible (we require IHL=5).
func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) { func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := newTestUDPCoalescer(t, w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv4(1000, 53, make([]byte, 200)) pkt := buildUDPv4(1000, 53, make([]byte, 200))
pkt[0] = 0x46 // IHL = 6 (24-byte IPv4 header — has options) pkt[0] = 0x46 // IHL = 6 (24-byte IPv4 header — has options)
if err := c.Commit(pkt); err != nil { if err := c.Commit(pkt); err != nil {
@@ -470,3 +479,58 @@ func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
t.Fatalf("ipv4-with-options must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites)) 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)
}
}
}
+47 -5
View File
@@ -8,10 +8,40 @@ import (
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum" 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 // TestChecksumMatchesGvisor walks lengths from 0 to 4096, with several initial
// seeds and a handful of starting alignments, asserting that our local // seeds and a handful of starting alignments, asserting that each local
// Checksum matches gvisor's reference bit-for-bit. // implementation matches gvisor's reference bit-for-bit.
func TestChecksumMatchesGvisor(t *testing.T) { 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)) rng := rand.New(rand.NewPCG(1, 2))
const padFront = 16 const padFront = 16
@@ -32,7 +62,7 @@ func TestChecksumMatchesGvisor(t *testing.T) {
} }
buf := pool[off : off+length] buf := pool[off : off+length]
want := gvisorchecksum.Checksum(buf, seed) want := gvisorchecksum.Checksum(buf, seed)
got := Checksum(buf, seed) got := impl.fn(buf, seed)
if got != want { if got != want {
t.Fatalf("len=%d off=%d seed=%#x: got %#04x want %#04x", t.Fatalf("len=%d off=%d seed=%#x: got %#04x want %#04x",
length, off, seed, got, want) length, off, seed, got, want)
@@ -40,12 +70,17 @@ func TestChecksumMatchesGvisor(t *testing.T) {
} }
} }
} }
})
}
} }
// TestChecksumPatternedBuffers exercises specific byte patterns that have // TestChecksumPatternedBuffers exercises specific byte patterns that have
// historically tripped up checksum implementations: all-zero, all-0xff, // historically tripped up checksum implementations: all-zero, all-0xff,
// alternating, and ascending sequences. // alternating, and ascending sequences.
func TestChecksumPatternedBuffers(t *testing.T) { 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++ { for length := 0; length <= 256; length++ {
patterns := map[string][]byte{ patterns := map[string][]byte{
"zeros": make([]byte, length), "zeros": make([]byte, length),
@@ -56,7 +91,7 @@ func TestChecksumPatternedBuffers(t *testing.T) {
for name, buf := range patterns { for name, buf := range patterns {
for _, seed := range []uint16{0, 0xffff, 0x8000} { for _, seed := range []uint16{0, 0xffff, 0x8000} {
want := gvisorchecksum.Checksum(buf, seed) want := gvisorchecksum.Checksum(buf, seed)
got := Checksum(buf, seed) got := impl.fn(buf, seed)
if got != want { if got != want {
t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x", t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x",
name, length, seed, got, want) name, length, seed, got, want)
@@ -64,6 +99,8 @@ func TestChecksumPatternedBuffers(t *testing.T) {
} }
} }
} }
})
}
} }
func bytes(n int, v byte) []byte { func bytes(n int, v byte) []byte {
@@ -98,6 +135,9 @@ func ascending(n int) []byte {
// and k=1 (one main loop iter, then tail). It's explicit coverage for // 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. // payload sizes that are odd, not divisible by 4, by 8, or by 32.
func TestChecksumTailPaths(t *testing.T) { 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)) rng := rand.New(rand.NewPCG(42, 17))
const padFront = 16 const padFront = 16
const maxK = 8 const maxK = 8
@@ -120,7 +160,7 @@ func TestChecksumTailPaths(t *testing.T) {
} }
buf := pool[off : off+length] buf := pool[off : off+length]
want := gvisorchecksum.Checksum(buf, seed) want := gvisorchecksum.Checksum(buf, seed)
got := Checksum(buf, seed) got := impl.fn(buf, seed)
if got != want { if got != want {
t.Fatalf("k=%d tail=%d (len=%d) off=%d seed=%#x: got %#04x want %#04x", t.Fatalf("k=%d tail=%d (len=%d) off=%d seed=%#x: got %#04x want %#04x",
k, tail, length, off, seed, got, want) k, tail, length, off, seed, got, want)
@@ -129,6 +169,8 @@ func TestChecksumTailPaths(t *testing.T) {
} }
} }
} }
})
}
} }
// BenchmarkChecksumTailSizes covers payload sizes that aren't clean multiples // BenchmarkChecksumTailSizes covers payload sizes that aren't clean multiples
+11
View File
@@ -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},
}
+8
View File
@@ -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},
}
+7
View File
@@ -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
+3 -4
View File
@@ -9,10 +9,9 @@ import (
"golang.org/x/sys/unix" "golang.org/x/sys/unix"
) )
// blockOn parks the calling goroutine until fd is ready (events is POLLIN for // blockOn parks the calling goroutine until fd is ready or shutdownFd signals teardown.
// reads, POLLOUT for writes) or shutdownFd signals teardown. It builds the // (events is POLLIN for reads, POLLOUT for writes)
// pollfd array on the stack every call, so concurrent callers on the same // It builds the pollfd array on the stack every call, so concurrent callers on the same Queue never share Revents storage.
// Queue never share Revents storage.
// //
// Returns os.ErrClosed when shutdown was signaled (POLLIN on shutdownFd) // Returns os.ErrClosed when shutdown was signaled (POLLIN on shutdownFd)
// or either fd reported a problem condition (POLLHUP|POLLNVAL|POLLERR). // or either fd reported a problem condition (POLLHUP|POLLNVAL|POLLERR).
+17 -15
View File
@@ -7,6 +7,7 @@ import (
"encoding/binary" "encoding/binary"
"errors" "errors"
"fmt" "fmt"
"log/slog"
"sync/atomic" "sync/atomic"
"golang.org/x/sys/unix" "golang.org/x/sys/unix"
@@ -17,18 +18,17 @@ type offloadQueueSet struct {
// pqi is exactly the same as pq, but stored as the interface type // pqi is exactly the same as pq, but stored as the interface type
pqi []Queue pqi []Queue
shutdownFd int shutdownFd int
// usoEnabled is true when newTun successfully negotiated TUN_F_USO4|6 // usoEnabled is true when newTun successfully negotiated TUN_F_USO4|6 with the kernel.
// with the kernel. Queues created by Add inherit this and surface it // Queues created by Add inherit this and surface it via Offload.USOSupported so coalescers can gate USO emission.
// via Offload.USOSupported so coalescers can gate USO emission.
usoEnabled bool usoEnabled bool
closed atomic.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 // NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do TSO segmentation.
// TSO segmentation in userspace. usoEnabled tells downstream queues whether // usoEnabled tells downstream queues whether the kernel agreed to deliver/accept GSO_UDP_L4 superpackets.
// the kernel agreed to deliver/accept GSO_UDP_L4 superpackets — coalescers func NewOffloadQueueSet(usoEnabled bool, l *slog.Logger) (QueueSet, error) {
// should fall back to per-packet writes when this is false.
func NewOffloadQueueSet(usoEnabled bool) (QueueSet, error) {
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC) shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to create eventfd: %w", err) return nil, fmt.Errorf("failed to create eventfd: %w", err)
@@ -39,6 +39,7 @@ func NewOffloadQueueSet(usoEnabled bool) (QueueSet, error) {
pqi: []Queue{}, pqi: []Queue{},
shutdownFd: shutdownFd, shutdownFd: shutdownFd,
usoEnabled: usoEnabled, usoEnabled: usoEnabled,
l: l,
} }
return out, nil return out, nil
@@ -49,7 +50,10 @@ func (c *offloadQueueSet) Queues() []Queue {
} }
func (c *offloadQueueSet) Add(fd int) error { func (c *offloadQueueSet) Add(fd int) error {
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled) if c.closed.Load() {
return errors.New("queue set already closed")
}
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled, c.l)
if err != nil { if err != nil {
return err return err
} }
@@ -73,23 +77,21 @@ func (c *offloadQueueSet) Close() error {
errs := []error{} errs := []error{}
// Signal all readers blocked in poll to wake up and exit. They observe // Signal all readers blocked in poll to wake up and exit.
// POLLIN on the shutdown eventfd and return os.ErrClosed. // They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
if err := c.wakeForShutdown(); err != nil { if err := c.wakeForShutdown(); err != nil {
errs = append(errs, err) errs = append(errs, err)
} }
// Close the per-queue tun fds; this also unblocks any in-flight reads. // Close the per-queue tun fds; this also unblocks any in-flight reads.
// The per-queue Close deliberately leaves shutdownFd alone - it belongs
// to this container.
for _, x := range c.pq { for _, x := range c.pq {
if err := x.Close(); err != nil { if err := x.Close(); err != nil {
errs = append(errs, err) errs = append(errs, err)
} }
} }
// Close the shutdown eventfd last: every reader's pollfd set references // Close the shutdown eventfd last: every reader's pollfd set references it,
// it, so it must outlive the wake + per-queue teardown above. // so it must outlive the wake + per-queue teardown above.
if err := unix.Close(c.shutdownFd); err != nil { if err := unix.Close(c.shutdownFd); err != nil {
errs = append(errs, err) errs = append(errs, err)
} }
+7 -6
View File
@@ -40,6 +40,9 @@ func (c *pollQueueSet) Queues() []Queue {
} }
func (c *pollQueueSet) Add(fd int) error { func (c *pollQueueSet) Add(fd int) error {
if c.closed.Load() {
return errors.New("queue set already closed")
}
x, err := newPoll(fd, c.shutdownFd) x, err := newPoll(fd, c.shutdownFd)
if err != nil { if err != nil {
return err return err
@@ -64,23 +67,21 @@ func (c *pollQueueSet) Close() error {
errs := []error{} errs := []error{}
// Wake any reader blocked in poll so it observes POLLIN on the shutdown // Signal all readers blocked in poll to wake up and exit.
// eventfd and returns os.ErrClosed. // They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
if err := c.wakeForShutdown(); err != nil { if err := c.wakeForShutdown(); err != nil {
errs = append(errs, err) errs = append(errs, err)
} }
// Close the per-queue tun fds; this also unblocks any in-flight reads. // Close the per-queue tun fds; this also unblocks any in-flight reads.
// The per-queue Close deliberately leaves shutdownFd alone - it belongs
// to this container.
for _, x := range c.pq { for _, x := range c.pq {
if err := x.Close(); err != nil { if err := x.Close(); err != nil {
errs = append(errs, err) errs = append(errs, err)
} }
} }
// Close the shutdown eventfd last: every reader's pollfd set references // Close the shutdown eventfd last: every reader's pollfd set references it,
// it, so it must outlive the wake + per-queue teardown above. // so it must outlive the wake + per-queue teardown above.
if err := unix.Close(c.shutdownFd); err != nil { if err := unix.Close(c.shutdownFd); err != nil {
errs = append(errs, err) errs = append(errs, err)
} }
+1 -7
View File
@@ -1,4 +1,4 @@
//go:build !linux || android || e2e_testing //go:build !linux || android
package tio package tio
@@ -8,12 +8,6 @@ func protoFromGSOType(_ uint8) (GSOProto, error) {
return 0, fmt.Errorf("GSO unsupported") return 0, fmt.Errorf("GSO unsupported")
} }
// SegmentSuperpacket invokes fn once per segment of pkt. On non-Linux
// builds (and Android/e2e_testing) this package does not provide a Queue
// implementation, so any caller that does construct a Packet here can only
// be operating on non-superpacket bytes and the stub forwards them
// directly. A non-zero GSO field is a programming error from the caller
// and returns an explicit error rather than silently misbehaving.
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error { func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
if pkt.GSO.IsSuperpacket() { if pkt.GSO.IsSuperpacket() {
return fmt.Errorf("tio: GSO superpacket on platform without segmentation support") return fmt.Errorf("tio: GSO superpacket on platform without segmentation support")
+5 -6
View File
@@ -4,9 +4,8 @@ import "io"
// singleQueue adapts a legacy one-datagram-per-Read source into a Queue. // singleQueue adapts a legacy one-datagram-per-Read source into a Queue.
// Read fills a private scratch buffer and returns exactly one Packet whose // 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 // Bytes borrow from that buffer, valid only until the next Read, per the Queue contract.
// Queue contract. Single-reader like every Queue; Write is exactly as safe // Single-reader like every Queue; Write is exactly as safe for concurrent use as the underlying source's Write.
// for concurrent use as the underlying source's Write.
type singleQueue struct { type singleQueue struct {
rw io.ReadWriter rw io.ReadWriter
closer io.Closer // nil: Close is a no-op (the source is shared and owned elsewhere) closer io.Closer // nil: Close is a no-op (the source is shared and owned elsewhere)
@@ -14,9 +13,9 @@ type singleQueue struct {
ret [1]Packet ret [1]Packet
} }
// NewSingleQueue wraps a one-datagram-per-Read ReadWriteCloser (a legacy tun // NewSingleQueue wraps a one-datagram-per-Read ReadWriteCloser (a legacy tun device) into a Queue.
// device) into a Queue. bufSize is the per-queue read scratch size and must // bufSize is the per-queue read scratch size and must be at least the largest datagram the source can return.
// be at least the largest datagram the source can return. Close closes rwc. // Close closes rwc.
func NewSingleQueue(rwc io.ReadWriteCloser, bufSize int) Queue { func NewSingleQueue(rwc io.ReadWriteCloser, bufSize int) Queue {
return &singleQueue{rw: rwc, closer: rwc, buf: make([]byte, bufSize)} return &singleQueue{rw: rwc, closer: rwc, buf: make([]byte, bufSize)}
} }
+60 -85
View File
@@ -13,79 +13,72 @@ type QueueSet interface {
Add(fd int) error Add(fd int) error
} }
// Capabilities advertises which kernel offload features a Queue // Capabilities advertises which kernel offload features a Queue successfully negotiated.
// successfully negotiated. Callers consult this to decide which coalescers // Callers consult this to decide which coalescers to wire onto the write path.
// to wire onto the write path — a Queue without TSO can't usefully accept a
// TCPCoalescer, and a Queue without USO can't accept a UDPCoalescer.
type Capabilities struct { type Capabilities struct {
// TSO means the FD was opened with IFF_VNET_HDR and the kernel agreed // TSO means the FD was opened with IFF_VNET_HDR and the kernel agreed to TUN_F_TSO4|TSO6,
// to TUN_F_TSO4|TSO6 — i.e. WriteGSO with GSOProtoTCP is safe. // and WriteGSO with GSOProtoTCP is safe.
TSO bool TSO bool
// USO means the kernel additionally agreed to TUN_F_USO4|USO6, so // USO means the kernel additionally agreed to TUN_F_USO4|USO6,
// WriteGSO with GSOProtoUDP is safe. Linux ≥ 6.2. // so WriteGSO with GSOProtoUDP is safe. Linux ≥ 6.2.
USO bool USO bool
} }
// Queue is a readable/writable Poll queue. Concurrency contract: a single // Queue is a readable/writable Poll queue.
// read goroutine drives Read; plain Write is safe for concurrent callers; // 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. // 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 { type Queue interface {
io.Closer io.Closer
// Read returns one or more packets. The returned Packet.Bytes slices // Read returns one or more packets.
// are borrowed from the Queue's internal buffer and are only valid // The returned Packet.Bytes slices are borrowed from the Queue's internal buffer and are only valid
// until the next Read or Close on this Queue - callers must encrypt // until the next Read or Close on this Queue.
// or copy each slice before the next call. A Packet may carry a // A Packet may carry a GSO/USO superpacket (see GSOInfo)
// GSO/USO superpacket (see GSOInfo); when GSO.IsSuperpacket() is // Single-reader only: not safe for concurrent Reads (it reuses per-queue rx scratch each call).
// true the caller must segment Bytes before treating it as a single
// IP datagram. Single-reader only: not safe for concurrent Reads (it
// reuses per-queue rx scratch each call).
Read() ([]Packet, error) Read() ([]Packet, error)
// Write emits a single packet on the plaintext (outside→inside) // Write emits a single packet on the plaintext (outside→inside) delivery path.
// delivery path. Safe for concurrent use. // Safe for concurrent use.
Write(p []byte) (int, error) Write(p []byte) (int, error)
} }
// Packet is the unit Queue.Read returns. Bytes points into the queue's // Packet is the unit Queue.Read returns.
// internal buffer and is only valid until the next Read or Close on the // Bytes points into the queue's internal buffer and is only valid until the next Read or Close on the queue that produced it.
// queue that produced it. GSO is the zero value for an already-segmented // GSO is the zero value for an already-segmented IP datagram;
// IP datagram; when non-zero it describes a kernel-supplied TSO/USO // when non-zero it describes a kernel-supplied TSO/USO superpacket the caller must segment before consuming.
// superpacket the caller must segment before consuming.
type Packet struct { type Packet struct {
Bytes []byte Bytes []byte
GSO GSOInfo GSO GSOInfo
} }
// GSOInfo describes a kernel-supplied superpacket sitting in Packet.Bytes. // GSOInfo describes a kernel-supplied superpacket sitting in Packet.Bytes.
// The zero value means "not a superpacket" — Bytes is one regular IP // The zero value means Bytes is one regular IP datagram and no segmentation is required.
// datagram and no segmentation is required.
type GSOInfo struct { type GSOInfo struct {
// Size is the GSO segment size: max payload bytes per segment // Size is the GSO segment size: max payload bytes per segment
// (== TCP MSS for TSO, == UDP payload chunk for USO). Zero means // (== TCP MSS for TSO, == UDP payload chunk for USO). Zero means not a superpacket.
// not a superpacket.
Size uint16 Size uint16
// HdrLen is the total L3+L4 header length within Bytes (already // HdrLen is the total L3+L4 header length within Bytes (already corrected via correctHdrLen, so safe to slice on).
// corrected via correctHdrLen, so safe to slice on).
HdrLen uint16 HdrLen uint16
// CsumStart is the L4 header offset inside Bytes (== L3 header // CsumStart is the L4 header offset inside Bytes (== L3 header length).
// length).
CsumStart uint16 CsumStart uint16
// Proto picks the L4 protocol (TCP or UDP) so the segmenter knows // Proto picks the L4 protocol (TCP or UDP) so the segmenter knows which checksum/header layout to apply.
// which checksum/header layout to apply.
Proto GSOProto Proto GSOProto
} }
// IsSuperpacket reports whether g describes a multi-segment GSO/USO // IsSuperpacket reports whether g describes a multi-segment GSO/USO
// superpacket that needs segmentation before its bytes can be encrypted // superpacket that needs segmentation before its bytes can be encrypted and sent on the wire.
// and sent on the wire.
func (g GSOInfo) IsSuperpacket() bool { return g.Size > 0 } func (g GSOInfo) IsSuperpacket() bool { return g.Size > 0 }
// Clone returns a Packet whose Bytes is a freshly allocated copy of p.Bytes, // Clone returns a Packet whose Bytes is a freshly allocated copy of p.Bytes,
// safe to retain past the next Read or Close on the originating Queue. // safe to retain past the next Read or Close on the originating Queue.
// GSO metadata is copied verbatim. Use this only when a caller genuinely // GSO metadata is copied verbatim.
// needs to outlive the borrowed-slice contract — the hot path reads should // Use this only when a caller needs the data to outlive the borrowed-slice contract.
// continue to consume the borrow synchronously to avoid the allocation.
func (p Packet) Clone() Packet { func (p Packet) Clone() Packet {
if p.Bytes == nil { if p.Bytes == nil {
return p return p
@@ -95,78 +88,60 @@ func (p Packet) Clone() Packet {
return Packet{Bytes: cp, GSO: p.GSO} return Packet{Bytes: cp, GSO: p.GSO}
} }
// CapsProvider is an optional interface implemented by Queues that // CapsProvider is an optional interface implemented by Queues that negotiate kernel offload features at open time.
// successfully negotiated kernel offload features at open time. Callers // Callers pick a write-path coalescer based on the result.
// pick a write-path coalescer based on the result. Queues that don't // Queues that don't implement it are treated as having no offload capability.
// implement it are treated as having no offload capability — callers must
// fall back to plain per-packet writes.
type CapsProvider interface { type CapsProvider interface {
Capabilities() Capabilities Capabilities() Capabilities
} }
// QueueCapabilities returns q's negotiated offload capabilities, or the // GSOProto selects the L4 protocol for a GSO superpacket.
// zero value when q does not advertise any. // Determines which VIRTIO_NET_HDR_GSO_* type the writer stamps and which checksum offset
func QueueCapabilities(q Queue) Capabilities {
if cp, ok := q.(CapsProvider); ok {
return cp.Capabilities()
}
return Capabilities{}
}
// GSOProto selects the L4 protocol for a GSO superpacket. Determines which
// VIRTIO_NET_HDR_GSO_* type the writer stamps and which checksum offset
// inside the transport header virtio NEEDS_CSUM expects. // inside the transport header virtio NEEDS_CSUM expects.
type GSOProto uint8 type GSOProto uint8
const ( const (
GSOProtoTCP GSOProto = iota GSOProtoUnknown GSOProto = iota
GSOProtoTCP
GSOProtoUDP GSOProtoUDP
) )
// GSOWriter is implemented by Queues that can emit a TCP or UDP superpacket // GSOWriter is implemented by Queues that can emit a TCP or UDP superpacket
// assembled from a header prefix plus one or more borrowed payload // assembled from a header prefix plus one or more borrowed payload fragments,
// fragments, in a single vectored write (writev with a leading // in a single vectored write (writev with a leading virtio_net_hdr).
// virtio_net_hdr). This lets the coalescer avoid copying payload bytes // This lets the coalescer avoid copying payload bytes between the caller's decrypt buffer and the TUN.
// between the caller's decrypt buffer and the TUN. Backends without GSO // Backends without GSO support do not implement this interface and coalescing is skipped.
// support do not implement this interface and coalescing is skipped.
// //
// hdr contains the IPv4/IPv6 header prefix (mutable - callers will have // hdr contains the IPv4/IPv6 header prefix (mutable: callers will have filled in total length and IP csum).
// filled in total length and IP csum). transportHdr is the TCP or UDP // transportHdr is the TCP or UDP header
// header (mutable - the L4 checksum field must hold the pseudo-header // (mutable: the L4 checksum field must hold the pseudo-header partial, single-fold not inverted, per virtio NEEDS_CSUM semantics).
// partial, single-fold not inverted, per virtio NEEDS_CSUM semantics). // pays are non-overlapping payload fragments whose concatenation is the full superpacket payload.
// pays are non-overlapping payload fragments whose concatenation is the // They are read-only from the writer's perspective and must remain valid until the call returns.
// full superpacket payload; they are read-only from the writer's // Every segment in pays except possibly the last must be exactly the same size.
// perspective and must remain valid until the call returns. Every segment // proto picks the L4 protocol so the writer knows which gsoType / CsumOffset to set.
// in pays except possibly the last is exactly the same size. proto picks
// the L4 protocol so the writer knows which GSOType / CsumOffset to set.
// //
// Callers should also consult CapsProvider (via SupportsGSO or // Callers should also consult CapsProvider (via SupportsGSO) for the per-protocol negotiated capability:
// QueueCapabilities) for the per-protocol negotiated capability; an // USO may not have been negotiated even when TSO was.
// implementation of GSOWriter is necessary but not sufficient since USO
// may not have been negotiated even when TSO was.
type GSOWriter interface { type GSOWriter interface {
io.Writer
CapsProvider
WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error
} }
// SupportsGSO reports whether w implements GSOWriter and the underlying // SupportsGSO reports whether w implements GSOWriter and the underlying
// queue advertises the negotiated capability for `want`. A writer that // queue advertises the negotiated capability for `want`.
// implements GSOWriter but not CapsProvider is treated as permissive func SupportsGSO(w io.Writer, want GSOProto) (GSOWriter, bool) {
// (used by tests and fakes that don't negotiate).
func SupportsGSO(w any, want GSOProto) (GSOWriter, bool) {
gw, ok := w.(GSOWriter) gw, ok := w.(GSOWriter)
if !ok { if !ok {
return nil, false return nil, false
} }
cp, ok := w.(CapsProvider) caps := gw.Capabilities()
if !ok {
return gw, true
}
caps := cp.Capabilities()
switch want { switch want {
case GSOProtoTCP: case GSOProtoTCP:
return gw, caps.TSO return gw, caps.TSO
case GSOProtoUDP: case GSOProtoUDP:
return gw, caps.USO return gw, caps.USO
} default:
return gw, false return gw, false
}
} }
+152 -149
View File
@@ -4,6 +4,7 @@
package tio package tio
import ( import (
"context"
"fmt" "fmt"
"io" "io"
"log/slog" "log/slog"
@@ -17,71 +18,56 @@ import (
"github.com/slackhq/nebula/overlay/tio/virtio" "github.com/slackhq/nebula/overlay/tio/virtio"
) )
// tunRxBufSize is the per-Read worst-case footprint inside rxBuf: one const maxSuperpacketLen = 65535
// kernel-supplied packet body, which is at most ~64 KiB (tunReadBufSize).
// 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 // Segmentation happens at encrypt time on a per-routine MTU-sized scratch
// (see SegmentSuperpacket), so rxBuf only holds raw kernel-supplied bytes. // (see SegmentSuperpacket), so rxBuf only holds raw kernel-supplied bytes.
// We round up to give comfortable margin for the drain headroom check // We round up to give margin for the drain headroom check below.
// below.
const tunRxBufSize = 64 * 1024 const tunRxBufSize = 64 * 1024
// tunRxBufCap is the total size we allocate for the per-reader rx // tunRxBufCap is the total size we allocate for the per-reader rx buffer.
// buffer. With reads landing directly in rxBuf, each drain iteration // Each drain iteration consumes up to tunRxBufSize of headroom for the kernel-supplied bytes.
// 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,
// Sized to eight such iterations so a single poll wake can drain several // amortizing the wake and giving the sendmmsg planner longer same-destination runs.
// TSO/USO superpackets under bulk load, amortizing the wake and giving // Hold latency stays bounded because listenIn flushes its send batch incrementally rather than only at end-of-drain.
// 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 const tunRxBufCap = tunRxBufSize * 8
// tunDrainCap caps how many packets a single Read will accumulate via // tunDrainCap caps how many packets a single Read will accumulate via the post-wake drain loop.
// the post-wake drain loop. Sized to soak up a burst of small ACKs while // Sized to soak up a burst of small ACKs while bounding how much work a single caller holds before handing off.
// bounding how much work a single caller holds before handing off.
const tunDrainCap = 64 const tunDrainCap = 64
// gsoMaxIovs caps the iovec budget WriteGSO assembles per call: 3 fixed // gsoMaxIovs caps the iovec budget WriteGSO assembles per call:
// entries (virtio_net_hdr, IP hdr, transport hdr) plus up to gsoMaxIovs-3 // 3 fixed entries (virtio_net_hdr, IP hdr, transport hdr), plus up to gsoMaxIovs-3 payload fragments.
// payload fragments. Sized comfortably above the typical kernel GSO // Sized comfortably above the typical kernel GSO segment cap (Linux UDP_GRO is 64)
// segment cap (Linux UDP_GRO is 64) so realistic coalesced bursts never // so realistic coalesced bursts never touch the limit.
// touch the limit. iovecs are tiny (16 bytes), so the entire scratch is // iovecs are tiny (16 bytes), so the entire scratch is 4 KiB.
// 4 KiB — fine to keep resident on every queue. WriteGSO returns an error // WriteGSO returns an error rather than reallocating when a caller exceeds this budget.
// rather than reallocating when a caller exceeds this budget.
const gsoMaxIovs = 256 const gsoMaxIovs = 256
// validVnetHdr is the 10-byte virtio_net_hdr we prepend to every non-GSO TUN // validVnetHdr is the 10-byte virtio_net_hdr we prepend to every non-GSO TUN write.
// write. Only flag set is VIRTIO_NET_HDR_F_DATA_VALID, which marks the skb // Only flag set is VIRTIO_NET_HDR_F_DATA_VALID. Note the tun write path
// CHECKSUM_UNNECESSARY so the receiving network stack skips L4 checksum // (__virtio_net_hdr_to_skb) ignores this bit — only the virtio-net driver's RX
// verification. All packets that reach the plain Write paths already carry // helper honors it — so packets land CHECKSUM_NONE and the stack verifies the
// a valid L4 checksum (either supplied by a remote peer whose ciphertext we // L4 checksum anyway. What matters here is what the header does NOT say:
// AEAD-authenticated, produced by segmentTCPYield/segmentUDPYield during // no NEEDS_CSUM, so the kernel is never asked to finish a checksum.
// superpacket segmentation, or built locally by CreateRejectPacket), so // All packets that reach the plain Write paths already carry a valid L4 checksum.
// trusting them is safe.
var validVnetHdr = [virtio.Size]byte{unix.VIRTIO_NET_HDR_F_DATA_VALID} 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. // 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. // 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 { type Offload struct {
fd int fd int
shutdownFd int shutdownFd int
closed atomic.Bool
rxBuf []byte // backing store for kernel-handed packets read this drain
rxOff int // cursor into rxBuf for the current Read drain
pending []Packet // packets returned from the most recent Read
// readVnetScratch holds the 10-byte virtio_net_hdr split off the front of
// every TUN read via readv(2). Decoupling the header from the packet body
// lets us read the body directly into rxBuf at the current rxOff with
// no userspace copy on the GSO_NONE fast path.
readVnetScratch [virtio.Size]byte
// readIovs is the readv(2) iovec scratch wired once at construction —
// iovec[0] points at readVnetScratch; iovec[1].Base/Len is updated per
// read to address the current rxBuf slot.
readIovs [2]unix.Iovec
// usoEnabled records whether the kernel agreed to TUN_F_USO* on this FD, // 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. // so writers can decide whether emitting GSO_UDP_L4 superpackets is safe.
usoEnabled bool usoEnabled bool
closed atomic.Bool
// gsoHdrBuf is a per-queue 10-byte scratch for the virtio_net_hdr emitted // 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 // by WriteGSO. Kept separate from the read-only package-level validVnetHdr
@@ -92,9 +78,27 @@ type Offload struct {
// gsoMaxIovs at construction; never grown. WriteGSO returns an error // gsoMaxIovs at construction; never grown. WriteGSO returns an error
// (and drops the call) if a caller hands it more fragments than fit. // (and drops the call) if a caller hands it more fragments than fit.
gsoIovs []unix.Iovec 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) (*Offload, error) { func newOffload(fd int, shutdownFd int, usoEnabled bool, l *slog.Logger) (*Offload, error) {
if err := unix.SetNonblock(fd, true); err != nil { if err := unix.SetNonblock(fd, true); err != nil {
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err) return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
} }
@@ -104,6 +108,7 @@ func newOffload(fd int, shutdownFd int, usoEnabled bool) (*Offload, error) {
shutdownFd: shutdownFd, shutdownFd: shutdownFd,
usoEnabled: usoEnabled, usoEnabled: usoEnabled,
closed: atomic.Bool{}, closed: atomic.Bool{},
l: l,
rxBuf: make([]byte, tunRxBufCap), rxBuf: make([]byte, tunRxBufCap),
gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs), gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs),
@@ -128,26 +133,18 @@ func (r *Offload) blockOnWrite() error {
return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLOUT) return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLOUT)
} }
// readPacket issues a single readv(2) splitting the virtio_net_hdr off // readPacket issues a single readv(2), splitting the virtio_net_hdr off into readVnetScratch
// into readVnetScratch and reading the packet body directly into rxBuf at // and reading the packet body directly into rxBuf at the current rxOff.
// the current rxOff. Returns the body length (zero virtio header bytes, // Returns the body length (zero virtio header bytes, just the IP packet/superpacket).
// just the IP packet/superpacket). block controls whether EAGAIN is // block controls whether EAGAIN is retried via poll: the initial read of a drain blocks; subsequent drain reads do not.
// retried via poll: the initial read of a drain blocks; subsequent drain
// reads do not.
//
// The body iovec capacity is always tunReadBufSize; callers (the Read
// drain loop) gate entry on tunRxBufCap-rxOff >= tunRxBufSize, sized to
// hold one worst-case kernel-supplied packet body. Without that gate the
// body iovec could be smaller than the next inbound packet and the
// kernel would truncate.
func (r *Offload) readPacket(block bool) (int, error) { func (r *Offload) readPacket(block bool) (int, error) {
for { for {
r.readIovs[1].Base = &r.rxBuf[r.rxOff] r.readIovs[1].Base = &r.rxBuf[r.rxOff]
r.readIovs[1].SetLen(tunReadBufSize) 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))) n, _, errno := syscall.Syscall(unix.SYS_READV, uintptr(r.fd), uintptr(unsafe.Pointer(&r.readIovs[0])), uintptr(len(r.readIovs)))
if errno == 0 { if errno == 0 {
if int(n) < virtio.Size { if int(n) < virtio.Size {
return 0, io.ErrShortWrite return 0, fmt.Errorf("tun read shorter than virtio_net_hdr: %d bytes", n)
} }
return int(n) - virtio.Size, nil return int(n) - virtio.Size, nil
} }
@@ -170,29 +167,30 @@ func (r *Offload) readPacket(block bool) (int, error) {
} }
} }
// Read returns one or more packets from the tun. Each Packet either // Read returns one or more packets from the tun.
// carries a single ready-to-use IP datagram (GSO zero) or a TSO/USO // 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).
// superpacket plus the GSOInfo a caller needs to segment it (see // The first read blocks via poll; once the fd is known readable we drain additional packets non-blocking until:
// SegmentSuperpacket). The first read blocks via poll; once the fd is // - the kernel queue is empty (EAGAIN)
// known readable we drain additional packets non-blocking until the // - we've collected tunDrainCap packets,
// kernel queue is empty (EAGAIN), we've collected tunDrainCap packets, // - or we're out of rxBuf headroom.
// 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 // This amortizes the poll wake over bursts of small packets (e.g. TCP ACKs).
// into the Offload's internal buffer and are only valid until the next // Packet.Bytes slices point into the Offload's internal buffer and are only valid until the next Read or Close on this Queue.
// Read or Close on this Queue.
func (r *Offload) Read() ([]Packet, error) { func (r *Offload) Read() ([]Packet, error) {
r.pending = r.pending[:0] r.pending = r.pending[:0]
r.rxOff = 0 r.rxOff = 0
// Initial (blocking) read. Retry on decode errors so a single bad // Initial (blocking) read.
// packet does not stall the reader. // Retry on decode errors so a single bad packet does not stall the reader.
for { for {
n, err := r.readPacket(true) n, err := r.readPacket(true)
if err != nil { if err != nil {
return nil, err return nil, err
} }
if err := r.decodeRead(n); err != nil { if err := r.decodeRead(n); err != nil {
// Drop and read again — a bad packet should not kill the reader. // 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 continue
} }
break break
@@ -214,6 +212,7 @@ func (r *Offload) Read() ([]Packet, error) {
if err := r.decodeRead(n); err != nil { if err := r.decodeRead(n); err != nil {
// Drop this packet and stop the drain; we'd rather hand off // Drop this packet and stop the drain; we'd rather hand off
// what we have than keep spinning here. // what we have than keep spinning here.
r.logDroppedRead(err)
break break
} }
} }
@@ -221,13 +220,20 @@ func (r *Offload) Read() ([]Packet, error) {
return r.pending, nil return r.pending, nil
} }
// decodeRead processes the packet sitting in rxBuf at rxOff (length // logDroppedRead reports a tun packet dropped for a bad/unsupported virtio
// pktLen). The bytes stay in rxBuf — for GSO_NONE we slice them as a // header. Debug-gated so the happy path never pays for attribute assembly.
// regular IP datagram (running finishChecksum if NEEDS_CSUM is set); func (r *Offload) logDroppedRead(err error) {
// for TSO/USO superpackets we attach the corrected GSO metadata so the if r.l != nil && r.l.Enabled(context.Background(), slog.LevelDebug) {
// caller can segment lazily at encrypt time. rxOff advances past the r.l.Debug("dropping tun packet with bad virtio header", "error", err)
// kernel-supplied body and nothing else, since segmentation no longer }
// writes back into rxBuf. }
// 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 { func (r *Offload) decodeRead(pktLen int) error {
if pktLen <= 0 { if pktLen <= 0 {
return fmt.Errorf("short tun read: %d", pktLen) return fmt.Errorf("short tun read: %d", pktLen)
@@ -237,7 +243,7 @@ func (r *Offload) decodeRead(pktLen int) error {
body := r.rxBuf[r.rxOff : r.rxOff+pktLen] body := r.rxBuf[r.rxOff : r.rxOff+pktLen]
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_NONE { if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_NONE {
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 { if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
if err := virtio.FinishChecksum(body, hdr); err != nil { if err := virtio.FinishChecksum(body, hdr); err != nil {
return err return err
@@ -248,17 +254,13 @@ func (r *Offload) decodeRead(pktLen int) error {
return nil return nil
} }
// GSO superpacket: validate, fix the kernel-supplied HdrLen on the
// FORWARD path (CorrectHdrLen), pick the L4 protocol, and attach
// the metadata. The bytes stay in rxBuf untouched, segmentation
// happens in SegmentSuperpacket at encrypt time.
if err := virtio.CheckValid(body, hdr); err != nil { if err := virtio.CheckValid(body, hdr); err != nil {
return err return err
} }
if err := virtio.CorrectHdrLen(body, &hdr); err != nil { if err := virtio.CorrectHdrLen(body, &hdr); err != nil {
return err return err
} }
proto, err := protoFromGSOType(hdr.GSOType) proto, err := protoFromGSOType(hdr.GSOType())
if err != nil { if err != nil {
return err return err
} }
@@ -276,22 +278,16 @@ func (r *Offload) decodeRead(pktLen int) error {
} }
func (r *Offload) Write(buf []byte) (int, error) { func (r *Offload) Write(buf []byte) (int, error) {
if len(buf) == 0 {
return 0, nil
}
iovs := [2]unix.Iovec{ iovs := [2]unix.Iovec{
{Base: &validVnetHdr[0]}, {Base: &validVnetHdr[0]},
{Base: &buf[0]}, {Base: &buf[0]},
} }
iovs[0].SetLen(virtio.Size) iovs[0].SetLen(virtio.Size)
iovs[1].SetLen(len(buf)) iovs[1].SetLen(len(buf))
return r.writeWithScratch(buf, &iovs) return r.rawWrite(unsafe.Slice(&iovs[0], 2))
}
func (r *Offload) writeWithScratch(buf []byte, iovs *[2]unix.Iovec) (int, error) {
if len(buf) == 0 {
return 0, nil
}
iovs[1].Base = &buf[0]
iovs[1].SetLen(len(buf))
return r.rawWrite(unsafe.Slice(&iovs[0], len(iovs)))
} }
func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) { func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
@@ -321,57 +317,34 @@ func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
// Capabilities reports the offload features negotiated for this Queue. TSO // Capabilities reports the offload features negotiated for this Queue. TSO
// is always true for Offload (we only construct it on IFF_VNET_HDR FDs); // 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 // USO is true only when the kernel agreed to TUN_F_USO4|6 at open time (Linux ≥ 6.2).
// (Linux ≥ 6.2).
func (r *Offload) Capabilities() Capabilities { func (r *Offload) Capabilities() Capabilities {
return Capabilities{TSO: true, USO: r.usoEnabled} return Capabilities{TSO: true, USO: r.usoEnabled}
} }
func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error { func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error {
if len(hdr) == 0 || len(pays) == 0 || len(transportHdr) == 0 { if len(pays) == 0 {
// There are no payload fragments. There is nothing to send.
return nil return nil
} }
// L4 checksum offset inside transportHdr: TCP=16 (the `check` field after var csumOff uint16 // csumOff is the offset of the L4 checksum field in transportHdr
// seq/ack/dataoff/flags/window), UDP=6 (after sport/dport/length).
var csumOff uint16
switch proto { switch proto {
case GSOProtoUDP: case GSOProtoUDP:
csumOff = 6 csumOff = 6
default: case GSOProtoTCP:
csumOff = 16 csumOff = 16
}
vhdr := virtio.Hdr{
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
HdrLen: uint16(len(hdr) + len(transportHdr)),
GSOSize: uint16(len(pays[0])),
CsumStart: uint16(len(hdr)),
CsumOffset: csumOff,
}
if len(pays) > 1 {
ipVer := hdr[0] >> 4
switch {
case proto == GSOProtoUDP && (ipVer == 4 || ipVer == 6):
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_UDP_L4
case ipVer == 6:
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV6
case ipVer == 4:
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV4
default: default:
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE return fmt.Errorf("unknown GSO proto: %d", proto)
vhdr.GSOSize = 0
} }
} else { // Incorrect geometry must cause an error, not a silent drop.
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE // No sane packet should ever make it inside this branch.
vhdr.GSOSize = 0 if len(hdr) == 0 || len(transportHdr) < int(csumOff)+2 {
return fmt.Errorf("tio: WriteGSO header too short: ip=%d transport=%d (csum field at %d)", len(hdr), len(transportHdr), csumOff)
} }
vhdr.Encode(r.gsoHdrBuf[:]) // Make the iovec array: [virtio_hdr, hdr, transportHdr, pays...].
// The constructor attaches r.gsoIovs[0] to gsoHdrBuf. That entry does not change.
// Build the iovec array: [virtio_hdr, hdr, transportHdr, pays...]. r.gsoIovs[0] is
// wired to gsoHdrBuf at construction and never changes.
need := 3 + len(pays) need := 3 + len(pays)
if need > cap(r.gsoIovs) { if need > cap(r.gsoIovs) {
slog.Default().Warn("tio: WriteGSO iovec budget exceeded; dropping superpacket",
"need", need, "cap", cap(r.gsoIovs), "segments", len(pays))
return fmt.Errorf("tio: WriteGSO needs %d iovecs but cap is %d", need, cap(r.gsoIovs)) return fmt.Errorf("tio: WriteGSO needs %d iovecs but cap is %d", need, cap(r.gsoIovs))
} }
r.gsoIovs = r.gsoIovs[:need] r.gsoIovs = r.gsoIovs[:need]
@@ -379,22 +352,52 @@ func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto
r.gsoIovs[1].SetLen(len(hdr)) r.gsoIovs[1].SetLen(len(hdr))
r.gsoIovs[2].Base = &transportHdr[0] r.gsoIovs[2].Base = &transportHdr[0]
r.gsoIovs[2].SetLen(len(transportHdr)) r.gsoIovs[2].SetLen(len(transportHdr))
// Defense in depth: an empty payload fragment can't be a valid GSO
// segment and &p[0] would panic on it. Callers route zero-length segSize := len(pays[0])
// datagrams through the plain path (see UDPCoalescer.commitParsed), so total := len(hdr) + len(transportHdr)
// this should never fire, but skip empties rather than index into one. for i, p := range pays {
// `n` tracks where the next payload iovec lands, since skips make it
// drift from 3+i.
n := 3
for _, p := range pays {
if len(p) == 0 { if len(p) == 0 {
continue // 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)
} }
r.gsoIovs[n].Base = &p[0] total += len(p)
r.gsoIovs[n].SetLen(len(p)) r.gsoIovs[3+i].Base = &p[0]
n++ r.gsoIovs[3+i].SetLen(len(p))
} }
r.gsoIovs = r.gsoIovs[:n] // This check keeps `total` in the uint16 range. Anything larger would wrap around and cause trouble.
if total > maxSuperpacketLen {
return fmt.Errorf("tio: WriteGSO superpacket %dB exceeds %d", total, maxSuperpacketLen)
}
// A single segment ships as a plain checksummed packet (GSO_NONE, size 0).
// Multiple segments carry the real GSO type and segSize, which the loop
// above verified is the size of every fragment except possibly the last.
gsoType := uint8(unix.VIRTIO_NET_HDR_GSO_NONE)
if len(pays) > 1 {
gsoType = gsoTypeFromProto(proto, hdr[0]>>4)
if gsoType == unix.VIRTIO_NET_HDR_GSO_NONE {
// gsoTypeFromProto only yields GSO_NONE for a bogus IP version nibble.
// A multi-segment superpacket must carry a real GSO type, or the kernel would deliver it as a single jumbo packet.
return fmt.Errorf("tio: WriteGSO IP version %d is not GSO-capable", hdr[0]>>4)
}
}
var gsoSize uint16
if gsoType != unix.VIRTIO_NET_HDR_GSO_NONE {
gsoSize = uint16(segSize)
}
virtio.EncodeHeader(
r.gsoHdrBuf[:],
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
gsoType, /*gsoType*/
uint16(len(hdr)+len(transportHdr)), /*hdrLen*/
gsoSize, /*gsoSize*/
uint16(len(hdr)), /*csumStart*/
csumOff, /*csumOffset*/
)
_, err := r.rawWrite(r.gsoIovs) _, err := r.rawWrite(r.gsoIovs)
return err return err
@@ -405,10 +408,10 @@ func (r *Offload) Close() error {
return nil return nil
} }
//shutdownFd is owned by the container, so we should not close it // 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 // 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.
// 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
// It gets EBADF -> os.ErrClosed (or wakes via the shutdown eventfd's // poll is NOT woken by this close; only the QueueSet's shutdown eventfd wake does that (see Queue.Close docs).
// ppoll first). closed.Swap already guarantees we only close once. // closed.Swap already guarantees we only close once.
return unix.Close(r.fd) return unix.Close(r.fd)
} }
+16 -15
View File
@@ -11,11 +11,6 @@ import (
"golang.org/x/sys/unix" "golang.org/x/sys/unix"
) )
// Maximum size we accept for a single read from a TUN with IFF_VNET_HDR. A
// TSO superpacket can be up to 64KiB of payload plus a single L2/L3/L4 header
// prefix plus the virtio header.
const tunReadBufSize = 65535
type Poll struct { type Poll struct {
fd int fd int
shutdownFd int shutdownFd int
@@ -25,10 +20,10 @@ type Poll struct {
batchRet [1]Packet batchRet [1]Packet
} }
// newPoll wraps an existing tun fd. On failure it does NOT close fd: the // newPoll wraps an existing tun fd.
// caller owns fd and is the sole closer (see pollQueueSet.Add callers in // On failure it does NOT close fd: the caller owns fd and is the sole closer
// overlay/tun_linux.go, which unix.Close on Add error). This matches the // (see pollQueueSet.Add callers in overlay/tun_linux.go, which unix.Close on Add error).
// newOffload convention and keeps closes at exactly one on every path. // This matches the newOffload convention and keeps closes at exactly one on every path.
func newPoll(fd int, shutdownFd int) (*Poll, error) { func newPoll(fd int, shutdownFd int) (*Poll, error) {
if err := unix.SetNonblock(fd, true); err != nil { if err := unix.SetNonblock(fd, true); err != nil {
return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err) return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err)
@@ -37,7 +32,7 @@ func newPoll(fd int, shutdownFd int) (*Poll, error) {
out := &Poll{ out := &Poll{
fd: fd, fd: fd,
shutdownFd: shutdownFd, shutdownFd: shutdownFd,
readBuf: make([]byte, tunReadBufSize), readBuf: make([]byte, 65535), // largest possible size Linux permits
} }
return out, nil return out, nil
} }
@@ -52,6 +47,11 @@ func (t *Poll) blockOnWrite() error {
return blockOn(int32(t.fd), int32(t.shutdownFd), unix.POLLOUT) 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) { func (t *Poll) Read() ([]Packet, error) {
n, err := t.readOne(t.readBuf) n, err := t.readOne(t.readBuf)
if err != nil { if err != nil {
@@ -108,10 +108,11 @@ func (t *Poll) Close() error {
if t.closed.Swap(true) { if t.closed.Swap(true) {
return nil 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 // shutdownFd is owned by the container, so we should not close it
// loading it in readOne, and mutating the field would race that load. // Close the underlying fd but do NOT null t.fd: a reader may still be loading it in readOne, and mutating the field would race that load.
// It gets EBADF -> os.ErrClosed (or wakes via the shutdown eventfd's // That reader gets EBADF -> os.ErrClosed on its next syscall. A reader already parked in
// ppoll first). closed.Swap already guarantees we only close once. // 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) return unix.Close(t.fd)
} }
+2 -1
View File
@@ -5,6 +5,7 @@ package tio
import ( import (
"errors" "errors"
"log/slog"
"os" "os"
"sync" "sync"
"testing" "testing"
@@ -210,7 +211,7 @@ func TestPollQueueSet_Close_ClosesShutdownFd(t *testing.T) {
// TestOffloadQueueSet_Close_ClosesShutdownFd mirrors the poll regression test // TestOffloadQueueSet_Close_ClosesShutdownFd mirrors the poll regression test
// for the GSO/offload queueset. // for the GSO/offload queueset.
func TestOffloadQueueSet_Close_ClosesShutdownFd(t *testing.T) { func TestOffloadQueueSet_Close_ClosesShutdownFd(t *testing.T) {
qs, err := NewOffloadQueueSet(false) qs, err := NewOffloadQueueSet(false, slog.New(slog.DiscardHandler))
require.NoError(t, err) require.NoError(t, err)
c, ok := qs.(*offloadQueueSet) c, ok := qs.(*offloadQueueSet)
require.True(t, ok) require.True(t, ok)
+25 -16
View File
@@ -1,5 +1,5 @@
//go:build linux && !android && !e2e_testing //go:build linux && !android
// +build linux,!android,!e2e_testing // +build linux,!android
package tio package tio
@@ -11,11 +11,11 @@ import (
"github.com/slackhq/nebula/overlay/tio/virtio" "github.com/slackhq/nebula/overlay/tio/virtio"
) )
// protoFromGSOType maps a virtio_net_hdr GSOType to the GSOProto value the // protoFromGSOType maps a virtio_net_hdr gsoType to the GSOProto value the
// segment-time helpers use. Returns an error for GSO_NONE or any unknown // segment-time helpers use. Returns an error for GSO_NONE or any unknown
// value — the caller should only invoke this on a confirmed superpacket. // value. The caller should only invoke this on a confirmed superpacket.
func protoFromGSOType(t uint8) (GSOProto, error) { func protoFromGSOType(t uint8) (GSOProto, error) {
switch t { switch t &^ unix.VIRTIO_NET_HDR_GSO_ECN {
case unix.VIRTIO_NET_HDR_GSO_TCPV4, unix.VIRTIO_NET_HDR_GSO_TCPV6: case unix.VIRTIO_NET_HDR_GSO_TCPV4, unix.VIRTIO_NET_HDR_GSO_TCPV6:
return GSOProtoTCP, nil return GSOProtoTCP, nil
case unix.VIRTIO_NET_HDR_GSO_UDP_L4: case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
@@ -25,17 +25,26 @@ func protoFromGSOType(t uint8) (GSOProto, error) {
} }
} }
// SegmentSuperpacket invokes fn once per segment of pkt. For non-GSO pkts // gsoTypeFromProto is the reverse of protoFromGSOType
// fn is called once with pkt.Bytes (no segmentation, no copy). For GSO/USO func gsoTypeFromProto(proto GSOProto, ipVer uint8) uint8 {
// superpackets fn is called once per segment with a slice of pkt.Bytes switch {
// holding that segment's plaintext (a freshly-patched L3+L4 header sliced case proto == GSOProtoUDP && (ipVer == 4 || ipVer == 6):
// in front of the original payload chunk). The slide is destructive: pkt is return unix.VIRTIO_NET_HDR_GSO_UDP_L4
// consumed by this call and its bytes are in an undefined state when case ipVer == 6:
// SegmentSuperpacket returns. Callers must not retain pkt or any earlier return unix.VIRTIO_NET_HDR_GSO_TCPV6
// seg slice past fn's return for that segment. The scratch parameter is case ipVer == 4:
// unused on the destructive path and kept only for cross-platform return unix.VIRTIO_NET_HDR_GSO_TCPV4
// signature compatibility. Aborts and returns the first error from fn or default:
// from per-segment construction. 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 { func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
if !pkt.GSO.IsSuperpacket() { if !pkt.GSO.IsSuperpacket() {
return fn(pkt.Bytes) return fn(pkt.Bytes)
+382 -86
View File
@@ -19,6 +19,36 @@ import (
// worst-case 64 KiB superpacket plus replicated per-segment headers). // worst-case 64 KiB superpacket plus replicated per-segment headers).
const testSegScratchSize = 192 * 1024 const testSegScratchSize = 192 * 1024
// TestProtoFromGSOTypeMasksECN guards the CWR-superpacket drop bug: the
// kernel qualifies a TSO superpacket whose TCP header carries CWR with
// VIRTIO_NET_HDR_GSO_ECN (we negotiate TUN_F_TSO_ECN, so it WILL send
// them once ECN feedback flows), and the decoder must mask that bit
// rather than reject the packet as an unknown type.
func TestProtoFromGSOTypeMasksECN(t *testing.T) {
cases := []struct {
typ uint8
want GSOProto
}{
{unix.VIRTIO_NET_HDR_GSO_TCPV4, GSOProtoTCP},
{unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, GSOProtoTCP},
{unix.VIRTIO_NET_HDR_GSO_TCPV6, GSOProtoTCP},
{unix.VIRTIO_NET_HDR_GSO_TCPV6 | unix.VIRTIO_NET_HDR_GSO_ECN, GSOProtoTCP},
{unix.VIRTIO_NET_HDR_GSO_UDP_L4, GSOProtoUDP},
}
for _, c := range cases {
got, err := protoFromGSOType(c.typ)
if err != nil || got != c.want {
t.Errorf("protoFromGSOType(%#x) = (%v, %v), want (%v, nil)", c.typ, got, err, c.want)
}
}
if _, err := protoFromGSOType(unix.VIRTIO_NET_HDR_GSO_NONE); err == nil {
t.Error("GSO_NONE must still be rejected")
}
if _, err := protoFromGSOType(unix.VIRTIO_NET_HDR_GSO_ECN); err == nil {
t.Error("a bare ECN bit with no base type must still be rejected")
}
}
// verifyChecksum confirms that the one's-complement sum across `b`, seeded // verifyChecksum confirms that the one's-complement sum across `b`, seeded
// with a folded pseudo-header sum, equals all-ones (valid). // with a folded pseudo-header sum, equals all-ones (valid).
func verifyChecksum(b []byte, pseudo uint16) bool { func verifyChecksum(b []byte, pseudo uint16) bool {
@@ -33,7 +63,7 @@ func verifyChecksum(b []byte, pseudo uint16) bool {
// returns. Tests pre-set hdr.HdrLen correctly, so correctHdrLen is not // returns. Tests pre-set hdr.HdrLen correctly, so correctHdrLen is not
// invoked here. // invoked here.
func segmentForTest(pkt []byte, hdr virtio.Hdr, out *[][]byte, scratch []byte) error { func segmentForTest(pkt []byte, hdr virtio.Hdr, out *[][]byte, scratch []byte) error {
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_NONE { if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_NONE {
cp := append([]byte(nil), pkt...) cp := append([]byte(nil), pkt...)
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 { if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
if err := virtio.FinishChecksum(cp, hdr); err != nil { if err := virtio.FinishChecksum(cp, hdr); err != nil {
@@ -43,7 +73,7 @@ func segmentForTest(pkt []byte, hdr virtio.Hdr, out *[][]byte, scratch []byte) e
*out = append(*out, cp) *out = append(*out, cp)
return nil return nil
} }
proto, err := protoFromGSOType(hdr.GSOType) proto, err := protoFromGSOType(hdr.GSOType())
if err != nil { if err != nil {
return err return err
} }
@@ -110,15 +140,14 @@ func buildTSOv4(t *testing.T, payLen, mss int) ([]byte, virtio.Hdr) {
for i := 0; i < payLen; i++ { for i := 0; i < payLen; i++ {
pkt[ipLen+tcpLen+i] = byte(i & 0xff) pkt[ipLen+tcpLen+i] = byte(i & 0xff)
} }
return pkt, virtio.NewHeader(
return pkt, virtio.Hdr{ unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_TCPV4, /*gsoType*/
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4, uint16(ipLen+tcpLen), /*hdrLen*/
HdrLen: uint16(ipLen + tcpLen), uint16(mss), /*gsoSize*/
GSOSize: uint16(mss), uint16(ipLen), /*csumStart*/
CsumStart: uint16(ipLen), 16, /*csumOffset*/
CsumOffset: 16, )
}
} }
func TestSegmentTCPv4(t *testing.T) { func TestSegmentTCPv4(t *testing.T) {
@@ -232,14 +261,14 @@ func TestSegmentTCPv6(t *testing.T) {
pkt[ipLen+tcpLen+i] = byte(i) pkt[ipLen+tcpLen+i] = byte(i)
} }
hdr := virtio.Hdr{ hdr := virtio.NewHeader(
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV6, unix.VIRTIO_NET_HDR_GSO_TCPV6, /*gsoType*/
HdrLen: uint16(ipLen + tcpLen), uint16(ipLen+tcpLen), /*hdrLen*/
GSOSize: uint16(mss), uint16(mss), /*gsoSize*/
CsumStart: uint16(ipLen), uint16(ipLen), /*csumStart*/
CsumOffset: 16, 16, /*csumOffset*/
} )
scratch := make([]byte, testSegScratchSize) scratch := make([]byte, testSegScratchSize)
var out [][]byte var out [][]byte
@@ -281,7 +310,7 @@ func TestSegmentTCPv6(t *testing.T) {
func TestSegmentGSONonePassesThrough(t *testing.T) { func TestSegmentGSONonePassesThrough(t *testing.T) {
pkt, hdr := buildTSOv4(t, 100, 100) pkt, hdr := buildTSOv4(t, 100, 100)
hdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE hdr.SetGSOType(unix.VIRTIO_NET_HDR_GSO_NONE)
hdr.Flags = 0 // no NEEDS_CSUM, leave packet untouched hdr.Flags = 0 // no NEEDS_CSUM, leave packet untouched
scratch := make([]byte, testSegScratchSize) scratch := make([]byte, testSegScratchSize)
@@ -300,7 +329,7 @@ func TestSegmentGSONonePassesThrough(t *testing.T) {
// TestSegmentRejectsLegacyUDPGSO ensures the legacy GSO_UDP (UFO) marker is // TestSegmentRejectsLegacyUDPGSO ensures the legacy GSO_UDP (UFO) marker is
// still rejected; only modern GSO_UDP_L4 (USO) is supported. // still rejected; only modern GSO_UDP_L4 (USO) is supported.
func TestSegmentRejectsLegacyUDPGSO(t *testing.T) { func TestSegmentRejectsLegacyUDPGSO(t *testing.T) {
hdr := virtio.Hdr{GSOType: unix.VIRTIO_NET_HDR_GSO_UDP} hdr := virtio.NewHeader(0, unix.VIRTIO_NET_HDR_GSO_UDP, 0, 0, 0, 0)
var out [][]byte var out [][]byte
if err := segmentForTest(nil, hdr, &out, nil); err == nil { if err := segmentForTest(nil, hdr, &out, nil); err == nil {
t.Fatalf("expected rejection for legacy UDP GSO") t.Fatalf("expected rejection for legacy UDP GSO")
@@ -324,22 +353,26 @@ func buildUSOv4(t *testing.T, payLen, gsoSize int) ([]byte, virtio.Hdr) {
copy(pkt[12:16], []byte{10, 0, 0, 1}) copy(pkt[12:16], []byte{10, 0, 0, 1})
copy(pkt[16:20], []byte{10, 0, 0, 2}) copy(pkt[16:20], []byte{10, 0, 0, 2})
// UDP header (length + checksum filled in per segment by segmentUDPYield) // UDP header. The kernel hands us a USO superpacket whose length field
// covers the WHOLE superpacket; the segmenter overwrites it per segment.
// Populating it here matters: leaving it zero makes the base-checksum path
// that must exclude it untestable, since excluding zero is a no-op.
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
binary.BigEndian.PutUint16(pkt[24:26], uint16(udpLen+payLen)) // superpacket length
for i := 0; i < payLen; i++ { for i := 0; i < payLen; i++ {
pkt[ipLen+udpLen+i] = byte(i & 0xff) pkt[ipLen+udpLen+i] = byte(i & 0xff)
} }
return pkt, virtio.Hdr{ return pkt, virtio.NewHeader(
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4, unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
HdrLen: uint16(ipLen + udpLen), uint16(ipLen+udpLen), /*hdrLen*/
GSOSize: uint16(gsoSize), uint16(gsoSize), /*gsoSize*/
CsumStart: uint16(ipLen), uint16(ipLen), /*csumStart*/
CsumOffset: 6, 6, /*csumOffset*/
} )
} }
func TestSegmentUDPv4(t *testing.T) { func TestSegmentUDPv4(t *testing.T) {
@@ -364,11 +397,12 @@ func TestSegmentUDPv4(t *testing.T) {
if totalLen != uint16(28+gso) { if totalLen != uint16(28+gso) {
t.Errorf("seg %d: total_len=%d want %d", i, totalLen, 28+gso) t.Errorf("seg %d: total_len=%d want %d", i, totalLen, 28+gso)
} }
// kernel UDP-GSO does NOT bump the IPv4 ID across segments; every // Software UDP GSO bumps the IPv4 ID per segment exactly like TSO
// segment carries the same ID as the seed. // (inet_gso_segment's fixed-ID case is TCP-only); wireguard-go's
// gsoSplit increments unconditionally too.
id := binary.BigEndian.Uint16(seg[4:6]) id := binary.BigEndian.Uint16(seg[4:6])
if id != 0x4242 { if id != 0x4242+uint16(i) {
t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242) t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242+uint16(i))
} }
udpLen := binary.BigEndian.Uint16(seg[24:26]) udpLen := binary.BigEndian.Uint16(seg[24:26])
if udpLen != uint16(8+gso) { if udpLen != uint16(8+gso) {
@@ -436,19 +470,21 @@ func TestSegmentUDPv6(t *testing.T) {
binary.BigEndian.PutUint16(pkt[40:42], 12345) binary.BigEndian.PutUint16(pkt[40:42], 12345)
binary.BigEndian.PutUint16(pkt[42:44], 53) binary.BigEndian.PutUint16(pkt[42:44], 53)
// Superpacket-wide length, as the kernel supplies it; see buildUSOv4.
binary.BigEndian.PutUint16(pkt[44:46], uint16(udpLen+payLen))
for i := 0; i < payLen; i++ { for i := 0; i < payLen; i++ {
pkt[ipLen+udpLen+i] = byte(i) pkt[ipLen+udpLen+i] = byte(i)
} }
hdr := virtio.Hdr{ hdr := virtio.NewHeader(
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4, unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
HdrLen: uint16(ipLen + udpLen), uint16(ipLen+udpLen), /*hdrLen*/
GSOSize: uint16(gso), uint16(gso), /*gsoSize*/
CsumStart: uint16(ipLen), uint16(ipLen), /*csumStart*/
CsumOffset: 6, 6, /*csumOffset*/
} )
scratch := make([]byte, testSegScratchSize) scratch := make([]byte, testSegScratchSize)
var out [][]byte var out [][]byte
@@ -580,14 +616,14 @@ func BenchmarkSegmentTCPv4(b *testing.B) {
for i := 0; i < sz.payLen; i++ { for i := 0; i < sz.payLen; i++ {
pkt[ipLen+tcpLen+i] = byte(i) pkt[ipLen+tcpLen+i] = byte(i)
} }
hdr := virtio.Hdr{ hdr := virtio.NewHeader(
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4, unix.VIRTIO_NET_HDR_GSO_TCPV4, /*gsoType*/
HdrLen: uint16(ipLen + tcpLen), uint16(ipLen+tcpLen), /*hdrLen*/
GSOSize: uint16(sz.mss), uint16(sz.mss), /*gsoSize*/
CsumStart: uint16(ipLen), uint16(ipLen), /*csumStart*/
CsumOffset: 16, 16, /*csumOffset*/
} )
scratch := make([]byte, testSegScratchSize) scratch := make([]byte, testSegScratchSize)
out := make([][]byte, 0, 64) out := make([][]byte, 0, 64)
@@ -640,35 +676,72 @@ func TestTunFileWriteVnetHdrNoAlloc(t *testing.T) {
} }
} }
// TestWriteGSOSkipsEmptyPayloads is the defense-in-depth guard for the // TestSegmentSuperpacketNoAlloc pins the segmenters' zero-allocation
// zero-length UDP DoS: a payload fragment of length zero would make &p[0] // contract. Both SegmentTCP and SegmentUDP derive their per-superpacket
// panic (index-out-of-range) when building the iovec array. WriteGSO must // constants into fixed-size arrays (tmp/ipTmp/savedHdr) that must stay on
// skip empties instead. We write to /dev/null so the writev always succeeds // the stack, and both take a yield closure that must not escape. Any of
// synchronously; the point is simply that neither call panics. // those escaping turns one allocation into one-per-superpacket on the
func TestWriteGSOSkipsEmptyPayloads(t *testing.T) { // hottest path in the reader, which BenchmarkSegmentSuperpacketAllocsTSO
fd, err := unix.Open("/dev/null", os.O_WRONLY, 0) // reports but nothing fails on. This does.
//
// The yield closure here only touches captured scalars: appending segments
// to a slice would allocate in the test itself and mask the measurement.
func TestSegmentSuperpacketNoAlloc(t *testing.T) {
const mss = 1400
const numSeg = 8
cases := []struct {
name string
build func() ([]byte, virtio.Hdr)
}{
{"tso-v4", func() ([]byte, virtio.Hdr) { return buildTSOv4(t, mss*numSeg, mss) }},
{"uso-v4", func() ([]byte, virtio.Hdr) { return buildUSOv4(t, mss*numSeg, mss) }},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
master, hdr := tc.build()
proto, err := protoFromGSOType(hdr.GSOType())
if err != nil { if err != nil {
t.Fatalf("open /dev/null: %v", err) t.Fatalf("protoFromGSOType: %v", err)
} }
t.Cleanup(func() { _ = unix.Close(fd) }) work := make([]byte, len(master))
p := Packet{Bytes: work, GSO: GSOInfo{
Size: hdr.GSOSize,
HdrLen: hdr.HdrLen,
CsumStart: hdr.CsumStart,
Proto: proto,
}}
o := &Offload{fd: fd, gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs)} // Segmentation consumes its input destructively, so restore from
o.gsoIovs[0].Base = &o.gsoHdrBuf[0] // the master copy each run; copy(2) into an existing slice does
o.gsoIovs[0].SetLen(virtio.Size) // not allocate. seen/bytes keep the closure from being optimized
// away and double as a sanity check that work actually happened.
ipHdr := make([]byte, 20) var seen, bytes int
ipHdr[0] = 0x45 // IPv4, IHL 5 run := func() {
udpHdr := make([]byte, 8) copy(work, master)
seen, bytes = 0, 0
// Sole payload empty: exercises the all-empty skip (n stays at 3). if err := SegmentSuperpacket(p, func(seg []byte) error {
if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{{}}, GSOProtoUDP); err != nil { seen++
t.Fatalf("WriteGSO with a single empty payload: %v", err) bytes += len(seg)
return nil
}); err != nil {
t.Fatalf("SegmentSuperpacket: %v", err)
} }
// Empty mixed with a real fragment: exercises the index-drift skip so a }
// later non-empty payload still lands in the right iovec slot.
real := make([]byte, 1200) run() // warm up: absorb any one-time allocation elsewhere
if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{real, {}}, GSOProtoUDP); err != nil { if seen != numSeg {
t.Fatalf("WriteGSO with a trailing empty payload: %v", err) t.Fatalf("yielded %d segments, want %d", seen, numSeg)
}
if allocs := testing.AllocsPerRun(200, run); allocs != 0 {
t.Fatalf("SegmentSuperpacket allocated %.1f times per call, want 0", allocs)
}
if seen != numSeg || bytes == 0 {
t.Fatalf("post-measure sanity: seen=%d bytes=%d", seen, bytes)
}
})
} }
} }
@@ -721,9 +794,11 @@ func TestDecodeReadFitsMaxTSOAtDrainThreshold(t *testing.T) {
const ipv6HdrLen = 40 const ipv6HdrLen = 40
const tcpHdrLen = 20 const tcpHdrLen = 20
const headerLen = ipv6HdrLen + tcpHdrLen const headerLen = ipv6HdrLen + tcpHdrLen
// Maximum TUN read body. The tunReadBufSize cap on readv's body iovec // Maximum TUN read body at the drain threshold. readv bounds the body
// is what bounds the kernel's superpacket length. // iovec by the space actually left in rxBuf, and the drain gate keeps that
pktLen := tunReadBufSize // at >= tunRxBufSize, so that is the largest superpacket the kernel can
// hand back on the last permitted drain read.
pktLen := tunRxBufSize
payLen := pktLen - headerLen payLen := pktLen - headerLen
const targetSegs = 64 const targetSegs = 64
gsoSize := (payLen + targetSegs - 1) / targetSegs gsoSize := (payLen + targetSegs - 1) / targetSegs
@@ -745,14 +820,14 @@ func TestDecodeReadFitsMaxTSOAtDrainThreshold(t *testing.T) {
copy(o.rxBuf[o.rxOff:], pkt) copy(o.rxBuf[o.rxOff:], pkt)
// Encode the matching virtio_net_hdr. // Encode the matching virtio_net_hdr.
hdr := virtio.Hdr{ hdr := virtio.NewHeader(
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV6, unix.VIRTIO_NET_HDR_GSO_TCPV6, /*gsoType*/
HdrLen: uint16(headerLen), uint16(headerLen), /*hdrLen*/
GSOSize: uint16(gsoSize), uint16(gsoSize), /*gsoSize*/
CsumStart: uint16(ipv6HdrLen), uint16(ipv6HdrLen), /*csumStart*/
CsumOffset: 16, 16, /*csumOffset*/
} )
hdr.Encode(o.readVnetScratch[:]) hdr.Encode(o.readVnetScratch[:])
startRxOff := o.rxOff startRxOff := o.rxOff
@@ -824,3 +899,224 @@ func TestDecodeReadFitsMaxTSOAtDrainThreshold(t *testing.T) {
t.Fatalf("got %d segments, want %d", gotSegs, wantSegs) t.Fatalf("got %d segments, want %d", gotSegs, wantSegs)
} }
} }
// TestOffloadWriteZeroLength: a zero-length Write must be a no-op, not a
// panic. The guard used to live below the &buf[0] that tripped on it.
func TestOffloadWriteZeroLength(t *testing.T) {
tf := &Offload{fd: -1} // any write reaching the fd would fail loudly
for _, buf := range [][]byte{nil, {}} {
n, err := tf.Write(buf)
if n != 0 || err != nil {
t.Errorf("Write(len=0) = (%d, %v), want (0, nil)", n, err)
}
}
}
// TestWriteGSOSuperpacketGeometry decodes the vnet header the kernel would see for a multi-segment write:
// the GSO type must match the proto and IP version
// gso_size must be the per-segment size (the kernel rejects a superpacket with gso_size == 0),
// and the csum fields must point at the transport header's checksum slot.
// Write through a pipe so the bytes can be read back and decoded.
func TestWriteGSOSuperpacketGeometry(t *testing.T) {
var pfds [2]int
if err := unix.Pipe(pfds[:]); err != nil {
t.Fatalf("pipe: %v", err)
}
t.Cleanup(func() { unix.Close(pfds[0]); unix.Close(pfds[1]) })
o := &Offload{fd: pfds[1], gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs)}
o.gsoIovs[0].Base = &o.gsoHdrBuf[0]
o.gsoIovs[0].SetLen(virtio.Size)
ipHdr := make([]byte, 20)
ipHdr[0] = 0x45
udpHdr := make([]byte, 8)
seg := make([]byte, 1200)
if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{seg, seg}, GSOProtoUDP); err != nil {
t.Fatalf("WriteGSO: %v", err)
}
buf := make([]byte, virtio.Size+len(ipHdr)+len(udpHdr)+2*len(seg)+64)
n, err := unix.Read(pfds[0], buf)
if err != nil {
t.Fatalf("read pipe: %v", err)
}
var vhdr virtio.Hdr
vhdr.Decode(buf[:virtio.Size])
if vhdr.GSOType() != unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
t.Errorf("gsoType=%d want UDP_L4", vhdr.GSOType())
}
if vhdr.GSOSize != 1200 {
t.Errorf("GSOSize=%d want 1200 (per-segment size from pays[0])", vhdr.GSOSize)
}
if vhdr.HdrLen != uint16(len(ipHdr)+len(udpHdr)) {
t.Errorf("HdrLen=%d want %d", vhdr.HdrLen, len(ipHdr)+len(udpHdr))
}
if vhdr.CsumStart != uint16(len(ipHdr)) || vhdr.CsumOffset != 6 {
t.Errorf("csum start/offset = %d/%d want %d/6", vhdr.CsumStart, vhdr.CsumOffset, len(ipHdr))
}
if want := virtio.Size + len(ipHdr) + len(udpHdr) + 2*len(seg); n != want {
t.Errorf("wrote %d bytes want %d", n, want)
}
}
// TestWriteGSORejectsBadGeometry pins the length-check contracts
func TestWriteGSORejectsBadGeometry(t *testing.T) {
fd, err := unix.Open("/dev/null", os.O_WRONLY, 0)
if err != nil {
t.Fatalf("open /dev/null: %v", err)
}
t.Cleanup(func() { _ = unix.Close(fd) })
o := &Offload{fd: fd, gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs)}
o.gsoIovs[0].Base = &o.gsoHdrBuf[0]
o.gsoIovs[0].SetLen(virtio.Size)
ipHdr := make([]byte, 20)
ipHdr[0] = 0x45
udpHdr := make([]byte, 8)
tcpHdr := make([]byte, 20)
seg := make([]byte, 1200)
cases := []struct {
name string
hdr, thdr []byte
pays [][]byte
proto GSOProto
wantErr bool
}{
{"empty-ip-hdr-with-payload", nil, udpHdr, [][]byte{seg}, GSOProtoUDP, true},
{"udp-transport-too-short-for-csum", ipHdr, udpHdr[:6], [][]byte{seg}, GSOProtoUDP, true},
{"tcp-transport-too-short-for-csum", ipHdr, tcpHdr[:16], [][]byte{seg}, GSOProtoTCP, true},
{"superpacket-over-65535", ipHdr, tcpHdr, [][]byte{make([]byte, 40000), make([]byte, 40000)}, GSOProtoTCP, true},
{"sole-payload-empty", ipHdr, udpHdr, [][]byte{{}}, GSOProtoUDP, true},
{"leading-empty-fragment", ipHdr, udpHdr, [][]byte{{}, seg, seg}, GSOProtoUDP, true},
{"trailing-empty-fragment", ipHdr, tcpHdr, [][]byte{seg, {}}, GSOProtoTCP, true},
{"oversize-middle-fragment", ipHdr, udpHdr, [][]byte{seg, make([]byte, 1201), seg}, GSOProtoUDP, true},
{"undersize-middle-fragment", ipHdr, udpHdr, [][]byte{seg, make([]byte, 100), seg}, GSOProtoUDP, true},
{"oversize-last-fragment", ipHdr, tcpHdr, [][]byte{seg, make([]byte, 1201)}, GSOProtoTCP, true},
{"short-last-fragment-ok", ipHdr, udpHdr, [][]byte{seg, seg, make([]byte, 100)}, GSOProtoUDP, false},
{"multi-segment-bad-ip-version", []byte{0x05}, udpHdr, [][]byte{seg, seg}, GSOProtoUDP, true},
{"single-segment-bad-ip-version-ok", []byte{0x05}, udpHdr, [][]byte{seg}, GSOProtoUDP, false},
{"no-pays-noop", ipHdr, udpHdr, nil, GSOProtoUDP, false},
{"valid-udp", ipHdr, udpHdr, [][]byte{seg, seg}, GSOProtoUDP, false},
{"valid-tcp", ipHdr, tcpHdr, [][]byte{seg, seg}, GSOProtoTCP, false},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
err := o.WriteGSO(tc.hdr, tc.thdr, tc.pays, tc.proto)
if tc.wantErr && err == nil {
t.Errorf("WriteGSO = nil, want error")
}
if !tc.wantErr && err != nil {
t.Errorf("WriteGSO = %v, want nil", err)
}
})
}
}
// BenchmarkSegmentUDPv4 is the USO counterpart to BenchmarkSegmentTCPv4. The
// yield is a no-op so the measurement is segmentation plus checksum work only.
func BenchmarkSegmentUDPv4(b *testing.B) {
sizes := []struct {
name string
payLen int
gsoSize int
}{
{"64KiB_GSO1400", 64000, 1400},
{"16KiB_GSO1400", 16384, 1400},
{"4KiB_GSO1400", 4096, 1400},
}
for _, sz := range sizes {
b.Run(sz.name, func(b *testing.B) {
const ipLen = 20
const udpLen = 8
pkt := make([]byte, ipLen+udpLen+sz.payLen)
pkt[0] = 0x45
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+udpLen+sz.payLen))
binary.BigEndian.PutUint16(pkt[4:6], 0x4242)
pkt[8] = 64
pkt[9] = unix.IPPROTO_UDP
copy(pkt[12:16], []byte{10, 0, 0, 1})
copy(pkt[16:20], []byte{10, 0, 0, 2})
binary.BigEndian.PutUint16(pkt[20:22], 12345)
binary.BigEndian.PutUint16(pkt[22:24], 53)
binary.BigEndian.PutUint16(pkt[24:26], uint16(udpLen+sz.payLen))
for i := 0; i < sz.payLen; i++ {
pkt[ipLen+udpLen+i] = byte(i)
}
master := append([]byte(nil), pkt...)
work := make([]byte, len(pkt))
p := Packet{Bytes: work, GSO: GSOInfo{
Size: uint16(sz.gsoSize),
HdrLen: ipLen + udpLen,
CsumStart: ipLen,
Proto: GSOProtoUDP,
}}
b.SetBytes(int64(len(pkt)))
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
copy(work, master)
if err := SegmentSuperpacket(p, func(seg []byte) error { return nil }); err != nil {
b.Fatal(err)
}
}
})
}
}
// BenchmarkSegmentUDPv6 mirrors BenchmarkSegmentUDPv4 for IPv6, where the
// pseudo-header address sum is 32 bytes rather than 8.
func BenchmarkSegmentUDPv6(b *testing.B) {
sizes := []struct {
name string
payLen int
gsoSize int
}{
{"64KiB_GSO1400", 64000, 1400},
{"16KiB_GSO1400", 16384, 1400},
{"4KiB_GSO1400", 4096, 1400},
}
for _, sz := range sizes {
b.Run(sz.name, func(b *testing.B) {
const ipLen = 40
const udpLen = 8
pkt := make([]byte, ipLen+udpLen+sz.payLen)
pkt[0] = 0x60
binary.BigEndian.PutUint16(pkt[4:6], uint16(udpLen+sz.payLen))
pkt[6] = unix.IPPROTO_UDP
pkt[7] = 64
pkt[8], pkt[9], pkt[23] = 0xfe, 0x80, 1
pkt[24], pkt[25], pkt[39] = 0xfe, 0x80, 2
binary.BigEndian.PutUint16(pkt[40:42], 12345)
binary.BigEndian.PutUint16(pkt[42:44], 53)
binary.BigEndian.PutUint16(pkt[44:46], uint16(udpLen+sz.payLen))
for i := 0; i < sz.payLen; i++ {
pkt[ipLen+udpLen+i] = byte(i)
}
master := append([]byte(nil), pkt...)
work := make([]byte, len(pkt))
p := Packet{Bytes: work, GSO: GSOInfo{
Size: uint16(sz.gsoSize),
HdrLen: ipLen + udpLen,
CsumStart: ipLen,
Proto: GSOProtoUDP,
}}
b.SetBytes(int64(len(pkt)))
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
copy(work, master)
if err := SegmentSuperpacket(p, func(seg []byte) error { return nil }); err != nil {
b.Fatal(err)
}
}
})
}
}
+41 -9
View File
@@ -3,7 +3,11 @@
package virtio package virtio
import "encoding/binary" import (
"encoding/binary"
"golang.org/x/sys/unix"
)
// Size is the on-wire length of struct virtio_net_hdr the kernel // Size is the on-wire length of struct virtio_net_hdr the kernel
// prepends/expects on a TUN opened with IFF_VNET_HDR (TUNSETVNETHDRSZ // prepends/expects on a TUN opened with IFF_VNET_HDR (TUNSETVNETHDRSZ
@@ -13,31 +17,59 @@ const Size = 10
// Hdr is the Go view of the legacy virtio_net_hdr. // Hdr is the Go view of the legacy virtio_net_hdr.
type Hdr struct { type Hdr struct {
Flags uint8 Flags uint8
GSOType uint8 gsoType uint8 //private to avoid mistakes wrt the 0x80 VIRTIO_NET_HDR_GSO_ECN flag, ORed with the other "GSO types"
HdrLen uint16 HdrLen uint16
GSOSize uint16 GSOSize uint16
CsumStart uint16 CsumStart uint16
CsumOffset 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 // Decode reads a virtio_net_hdr in host byte order (TUN default; we never
// call TUNSETVNETLE so the kernel matches our endianness). // call TUNSETVNETLE so the kernel matches our endianness).
func (h *Hdr) Decode(b []byte) { func (h *Hdr) Decode(b []byte) {
h.Flags = b[0] h.Flags = b[0]
h.GSOType = b[1] h.gsoType = b[1]
h.HdrLen = binary.NativeEndian.Uint16(b[2:4]) h.HdrLen = binary.NativeEndian.Uint16(b[2:4])
h.GSOSize = binary.NativeEndian.Uint16(b[4:6]) h.GSOSize = binary.NativeEndian.Uint16(b[4:6])
h.CsumStart = binary.NativeEndian.Uint16(b[6:8]) h.CsumStart = binary.NativeEndian.Uint16(b[6:8])
h.CsumOffset = binary.NativeEndian.Uint16(b[8:10]) 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 // 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. // (must be at least Size bytes). Used to emit a TSO superpacket on egress.
func (h *Hdr) Encode(b []byte) { func (h *Hdr) Encode(b []byte) {
b[0] = h.Flags EncodeHeader(b, h.Flags, h.gsoType, h.HdrLen, h.GSOSize, h.CsumStart, h.CsumOffset)
b[1] = h.GSOType }
binary.NativeEndian.PutUint16(b[2:4], h.HdrLen)
binary.NativeEndian.PutUint16(b[4:6], h.GSOSize) // GSOType returns gsoType with the ECN-flag masked out
binary.NativeEndian.PutUint16(b[6:8], h.CsumStart) func (h *Hdr) GSOType() uint8 {
binary.NativeEndian.PutUint16(b[8:10], h.CsumOffset) 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
} }
+138 -131
View File
@@ -27,11 +27,8 @@ const (
tcpHeaderMaxLen = 60 // data-offset=15, max options tcpHeaderMaxLen = 60 // data-offset=15, max options
) )
// maxSegHdrLen bounds the L3+L4 header we snapshot before stamping each // maxSegHdrLen bounds the L3+L4 header we snapshot before stamping each segment.
// segment. The largest header the segmenter supports is IPv4 (max IHL 60) // The largest header the segmenter supports is IPv4 (max IHL 60) plus TCP (max data-offset 60) = 120 bytes
// plus TCP (max data-offset 60) = 120 bytes; the array is sized to that
// worst case so the snapshot lives on the stack with no per-call heap
// allocation.
const maxSegHdrLen = ipv4HeaderMaxLen + tcpHeaderMaxLen // 120 const maxSegHdrLen = ipv4HeaderMaxLen + tcpHeaderMaxLen // 120
// Byte offsets inside an IPv4 header. // Byte offsets inside an IPv4 header.
@@ -65,63 +62,72 @@ const (
udpChecksumOff = 6 udpChecksumOff = 6
) )
var errPacketTooShort = errors.New("packet too short")
// tcpFinPshMask is cleared on every segment except the last of a TSO burst. // tcpFinPshMask is cleared on every segment except the last of a TSO burst.
const tcpFinPshMask = 0x09 // FIN(0x01) | PSH(0x08) const tcpFinPshMask = 0x09 // FIN(0x01) | PSH(0x08)
// tcpCwrFlag is cleared on every segment except the first. Per RFC 3168 // tcpCwrFlag is cleared on every segment except the first.
// §6.1.2 the CWR bit signals a one-shot transition (the sender just halved // Per RFC 3168 §6.1.2 the CWR bit signals a one-shot transition (the sender just halved its window)
// its window) and must appear on the first segment of a TSO burst only. // and must appear on the first segment of a TSO burst only.
const tcpCwrFlag = 0x80 const tcpCwrFlag = 0x80
// CheckValid rejects packets whose virtio_net_hdr/IP combination would // CheckValid rejects packets whose virtio_net_hdr/IP combination would
// cause a downstream miscompute. The TUN should never emit RSC_INFO and // cause a downstream miscompute. The TUN should never emit RSC_INFO and
// the GSO type must agree with the IP version nibble. // the GSO type must agree with the IP version nibble.
func CheckValid(pkt []byte, hdr Hdr) error { func CheckValid(pkt []byte, hdr Hdr) error {
// When RSC_INFO is set the csum_start/csum_offset fields are repurposed to
// carry coalescing info rather than checksum offsets. A TUN writing via
// IFF_VNET_HDR should never emit this, but if it did we would silently
// miscompute the segment checksums — refuse the packet instead.
if hdr.Flags&unix.VIRTIO_NET_HDR_F_RSC_INFO != 0 { if hdr.Flags&unix.VIRTIO_NET_HDR_F_RSC_INFO != 0 {
return fmt.Errorf("virtio RSC_INFO flag not supported on TUN reads") return fmt.Errorf("virtio RSC_INFO flag not supported on TUN reads")
} }
if len(pkt) < ipv4HeaderMinLen { if len(pkt) < ipv4HeaderMinLen {
return fmt.Errorf("packet too short") return errPacketTooShort
} }
ipVersion := pkt[0] >> 4 ipVersion := pkt[0] >> 4
switch hdr.GSOType { 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: case unix.VIRTIO_NET_HDR_GSO_TCPV4:
if ipVersion != 4 { if ipVersion != 4 {
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType) return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
} }
case unix.VIRTIO_NET_HDR_GSO_TCPV6: case unix.VIRTIO_NET_HDR_GSO_TCPV6:
if ipVersion != 6 { if ipVersion != 6 {
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType) return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
} }
case unix.VIRTIO_NET_HDR_GSO_UDP_L4: case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
// USO carries either v4 or v6; the leading nibble disambiguates. // USO carries either v4 or v6; the leading nibble disambiguates.
if !(ipVersion == 4 || ipVersion == 6) { if !(ipVersion == 4 || ipVersion == 6) {
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType) return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
} }
default: default:
if !(ipVersion == 6 || ipVersion == 4) { if !(ipVersion == 6 || ipVersion == 4) {
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType) return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
} }
} }
return nil return nil
} }
// CorrectHdrLen rewrites hdr.HdrLen based on the actual transport header // CorrectHdrLen rewrites hdr.HdrLen based on the actual transport header length read out of pkt.
// length read out of pkt. The kernel's hdr.HdrLen on the FORWARD path can // The kernel's hdr.HdrLen on the FORWARD path can be the length of the entire first packet, so we don't trust it.
// be the length of the entire first packet, so we don't trust it.
func CorrectHdrLen(pkt []byte, hdr *Hdr) error { func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
// Thank you wireguard-go for documenting these edge-cases // 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 // 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 // of the entire first packet when the kernel is handling it as part of a FORWARD path.
// FORWARD path. Instead, parse the transport header length and add it onto // Instead, parse the transport header length and add it onto csumStart, which is synonymous for IP header length.
// csumStart, which is synonymous for IP header length.
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_UDP_L4 { if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
hdr.HdrLen = hdr.CsumStart + 8 hdr.HdrLen = hdr.CsumStart + 8
} else { } else {
if len(pkt) <= int(hdr.CsumStart+tcpDataOffOff) { if len(pkt) <= int(hdr.CsumStart+tcpDataOffOff) {
@@ -129,8 +135,7 @@ func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
} }
tcpHLen := uint16(pkt[hdr.CsumStart+tcpDataOffOff] >> 4 * 4) tcpHLen := uint16(pkt[hdr.CsumStart+tcpDataOffOff] >> 4 * 4)
if tcpHLen < 20 || tcpHLen > 60 { if tcpHLen < tcpHeaderMinLen || tcpHLen > tcpHeaderMaxLen {
// A TCP header must be between 20 and 60 bytes in length.
return fmt.Errorf("tcp header len is invalid: %d", tcpHLen) return fmt.Errorf("tcp header len is invalid: %d", tcpHLen)
} }
hdr.HdrLen = hdr.CsumStart + tcpHLen hdr.HdrLen = hdr.CsumStart + tcpHLen
@@ -150,19 +155,62 @@ func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
return nil return nil
} }
// SegmentTCP walks a TSO superpacket pkt, yielding each segment as a // segCount returns how many segments a payload of payLen bytes splits into at gsoSize,
// slice into pkt itself. Per-segment plaintext is laid out by stamping a // with a floor of one so a header-only superpacket still yields a single segment.
// copy of the original L3+L4 header into pkt at offset i*gsoSize, where it func segCount(payLen, gsoSize int) int {
// sits immediately before that segment's payload chunk in the original n := (payLen + gsoSize - 1) / gsoSize
// buffer. The stamp is destructive but harmless: iter i's header write lands if n == 0 {
// on pkt[i*G : i*G+hdrLen], which is the tail of seg_{i-1}'s payload (already return 1
// consumed) and ends exactly where seg_i's payload begins, so it never clobbers }
// live payload — this holds even when gsoSize < hdrLen. The header bytes are return n
// sourced from a pristine snapshot taken before the loop (savedHdr), NOT from }
// pkt[:hdrLen], because when gsoSize < hdrLen the stamps would otherwise
// overwrite the leading header in place and every stamp after the first would // basePseudoSum folds the part of the L4 pseudo-header sum that is identical
// copy corrupted bytes. pkt is consumed by this call and must not be inspected // for every segment: the source and destination addresses plus the protocol
// by the caller after the final yield. // 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 { func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
if gsoSizeU == 0 { if gsoSizeU == 0 {
return fmt.Errorf("gso_size is zero") return fmt.Errorf("gso_size is zero")
@@ -181,49 +229,28 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4 tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4
payLen := len(pkt) - headerLen payLen := len(pkt) - headerLen
gsoSize := int(gsoSizeU) gsoSize := int(gsoSizeU)
numSeg := (payLen + gsoSize - 1) / gsoSize numSeg := segCount(payLen, gsoSize)
if numSeg == 0 {
numSeg = 1
}
origSeq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4]) origSeq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
origFlags := pkt[csumStart+tcpFlagsOff] origFlags := pkt[csumStart+tcpFlagsOff]
var tmp [tcpHeaderMaxLen]byte baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_TCP)
copy(tmp[:tcpHdrLen], pkt[csumStart:headerLen]) baseTcpHdrSum := baseTCPHdrSum(pkt, csumStart, headerLen)
tmp[tcpSeqOff], tmp[tcpSeqOff+1], tmp[tcpSeqOff+2], tmp[tcpSeqOff+3] = 0, 0, 0, 0
tmp[tcpFlagsOff] = 0
tmp[tcpChecksumOff], tmp[tcpChecksumOff+1] = 0, 0
baseTcpHdrSum := uint32(checksum.Checksum(tmp[:tcpHdrLen], 0))
var baseProtoSum uint32
if isV4 {
baseProtoSum = uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0))
} else {
baseProtoSum = uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0))
}
baseProtoSum += uint32(unix.IPPROTO_TCP)
var origIPID uint16 var origIPID uint16
var baseIPHdrSum uint32 var baseIPHdrSum uint32
if isV4 { if isV4 {
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2]) origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
ihl := int(pkt[0]&0x0f) * 4 var err error
if ihl < ipv4HeaderMinLen || ihl > csumStart { // TSO bumps the ID per segment, so it stays out of the base sum.
return fmt.Errorf("bad IPv4 IHL: %d", ihl) baseIPHdrSum, err = baseIPv4HdrSum(pkt, csumStart)
if err != nil {
return err
} }
var ipTmp [ipv4HeaderMaxLen]byte
copy(ipTmp[:ihl], pkt[:ihl])
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
} }
// Snapshot the pristine L3+L4 header once. Every segment's header is // Snapshot the pristine L3+L4 header once. '
// stamped from this copy, so overlapping stamps (gsoSize < headerLen) // Every segment's header is stamped from this copy, so overlapping stamps (gsoSize < headerLen) can never corrupt the source.
// can never corrupt the source. The variable fields (seq/flags/cksum/
// totalLen/id) captured here are stale but are overwritten per segment.
var savedHdr [maxSegHdrLen]byte var savedHdr [maxSegHdrLen]byte
copy(savedHdr[:headerLen], pkt[:headerLen]) copy(savedHdr[:headerLen], pkt[:headerLen])
@@ -237,12 +264,10 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
segLen := headerLen + segPayLen segLen := headerLen + segPayLen
headerOff := i * gsoSize headerOff := i * gsoSize
// Stamp the header into place immediately before this segment's // Stamp the header into place immediately before this segment's payload, sourced from the snapshot.
// payload, sourced from the pristine snapshot. Iter 0's header is // The per-segment patches below overwrite the variable fields. (seq/flags/cksum/totalLen/id)
// already at pkt[:headerLen] (identical to savedHdr), so only i ≥ 1
// needs the stamp. The per-segment patches below overwrite the
// variable fields.
if i > 0 { 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]) copy(pkt[headerOff:headerOff+headerLen], savedHdr[:headerLen])
} }
seg := pkt[headerOff : headerOff+segLen] seg := pkt[headerOff : headerOff+segLen]
@@ -271,10 +296,9 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
seg[csumStart+tcpFlagsOff] = segFlags seg[csumStart+tcpFlagsOff] = segFlags
tcpLen := tcpHdrLen + segPayLen tcpLen := tcpHdrLen + segPayLen
// Payload bytes still live at their original offset in pkt. The // Payload bytes still live at their original offset in pkt.
// header slide above only writes into pkt[i*G : i*G+H], which is // 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)
// the tail of seg_{i-1}'s payload (already consumed) and never // and never overlaps seg_i's own payload at pkt[header+i*GSOSize : header+(i+1)*GSOSize].
// overlaps seg_i's own payload at pkt[H+i*G : H+(i+1)*G].
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0)) paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
wide := uint64(baseTcpHdrSum) + uint64(paySum) + uint64(baseProtoSum) wide := uint64(baseTcpHdrSum) + uint64(paySum) + uint64(baseProtoSum)
wide += uint64(segSeq) + uint64(segFlags) + uint64(tcpLen) wide += uint64(segSeq) + uint64(segFlags) + uint64(tcpLen)
@@ -290,17 +314,10 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
return nil return nil
} }
// SegmentUDP walks a USO superpacket, stamping a per-segment-patched copy of // SegmentUDP walks a USO superpacket, stamping a per-segment-patched copy of the original L3+L4 header
// the original L3+L4 header into pkt at offset i*gsoSize and yielding // into pkt at offset i*GSOSize and yielding pkt[i*GSOSize:i*GSOSize+segLen] to the caller.
// pkt[i*G:i*G+segLen] to the caller. Per-segment patches are total_len + // Per-segment patches are total_len + IPv4 csum (or IPv6 payload_len) plus the UDP length and checksum.
// IPv4 csum (or IPv6 payload_len) plus the UDP length and checksum. pkt is // pkt is consumed destructively.
// consumed destructively; see SegmentTCP for the layout reasoning, including
// why the header is stamped from a pristine snapshot rather than pkt[:hdrLen]
// (correctness when gsoSize < hdrLen).
//
// UDP-GSO leaves the IPv4 ID identical across segments (the kernel does not
// bump it), which is why the IP-level per-segment work is limited to
// total_len + IPv4 header checksum (v4) or payload_len (v6).
func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error { func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
if gsoSizeU == 0 { if gsoSizeU == 0 {
return fmt.Errorf("gso_size is zero") return fmt.Errorf("gso_size is zero")
@@ -321,41 +338,24 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
payLen := len(pkt) - headerLen payLen := len(pkt) - headerLen
gsoSize := int(gsoSizeU) gsoSize := int(gsoSizeU)
numSeg := (payLen + gsoSize - 1) / gsoSize numSeg := segCount(payLen, gsoSize)
if numSeg == 0 {
numSeg = 1
}
var udpTmp [udpHeaderLen]byte baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_UDP)
copy(udpTmp[:], pkt[csumStart:headerLen])
udpTmp[udpLengthOff], udpTmp[udpLengthOff+1] = 0, 0
udpTmp[udpChecksumOff], udpTmp[udpChecksumOff+1] = 0, 0
baseUDPHdrSum := uint32(checksum.Checksum(udpTmp[:], 0))
var baseProtoSum uint32
if isV4 {
baseProtoSum = uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0))
} else {
baseProtoSum = uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0))
}
baseProtoSum += uint32(unix.IPPROTO_UDP)
var origIPID uint16
var baseIPHdrSum uint32 var baseIPHdrSum uint32
if isV4 { if isV4 {
ihl := int(pkt[0]&0x0f) * 4 origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
if ihl < ipv4HeaderMinLen || ihl > csumStart { var err error
return fmt.Errorf("bad IPv4 IHL: %d", ihl) // 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
} }
var ipTmp [ipv4HeaderMaxLen]byte
copy(ipTmp[:ihl], pkt[:ihl])
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
} }
// Snapshot the pristine L3+L4 header once and stamp every segment from // Snapshot the pristine L3+L4 header once and stamp every segment from it
// it; see SegmentTCP for why sourcing from pkt[:headerLen] corrupts
// segments when gsoSize < headerLen.
var savedHdr [maxSegHdrLen]byte var savedHdr [maxSegHdrLen]byte
copy(savedHdr[:headerLen], pkt[:headerLen]) copy(savedHdr[:headerLen], pkt[:headerLen])
@@ -378,8 +378,10 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
udpLen := udpHeaderLen + segPayLen udpLen := udpHeaderLen + segPayLen
if isV4 { if isV4 {
segID := origIPID + uint16(i)
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen)) binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
ipSum := baseIPHdrSum + uint32(totalLen) binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID)
ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID)
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum)) binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
} else { } else {
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen)) binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
@@ -387,12 +389,13 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
binary.BigEndian.PutUint16(seg[csumStart+udpLengthOff:csumStart+udpLengthOff+2], uint16(udpLen)) binary.BigEndian.PutUint16(seg[csumStart+udpLengthOff:csumStart+udpLengthOff+2], uint16(udpLen))
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0)) // Sum the UDP header (length just written, checksum zeroed) together with
wide := uint64(baseUDPHdrSum) + uint64(paySum) + uint64(baseProtoSum) // this segment's payload in one pass, seeded with the pseudo-header sum.
wide += uint64(udpLen) + uint64(udpLen) seg[csumStart+udpChecksumOff], seg[csumStart+udpChecksumOff+1] = 0, 0
wide = (wide & 0xffffffff) + (wide >> 32) pseudo := baseProtoSum + uint32(udpLen)
wide = (wide & 0xffffffff) + (wide >> 32) pseudo = (pseudo & 0xffff) + (pseudo >> 16)
csum := foldComplement(uint32(wide)) pseudo = (pseudo & 0xffff) + (pseudo >> 16)
csum := ^checksum.Checksum(seg[csumStart:], uint16(pseudo))
if csum == 0 { if csum == 0 {
csum = 0xffff csum = 0xffff
} }
@@ -406,10 +409,9 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
return nil return nil
} }
// FinishChecksum computes the L4 checksum for a non-GSO packet that the kernel // FinishChecksum computes the L4 checksum for a non-GSO packet that the kernel handed us with NEEDS_CSUM set.
// handed us with NEEDS_CSUM set. csum_start / csum_offset point at the 16-bit // CsumStart / CsumOffset point at the 16-bit checksum field.
// checksum field; we zero it, fold a full sum (the field was pre-loaded with // We zero it, fold a full sum from the partial one that the kernel provided, and store the result.
// the pseudo-header partial sum by the kernel), and store the result.
func FinishChecksum(seg []byte, hdr Hdr) error { func FinishChecksum(seg []byte, hdr Hdr) error {
cs := int(hdr.CsumStart) cs := int(hdr.CsumStart)
co := int(hdr.CsumOffset) co := int(hdr.CsumOffset)
@@ -421,7 +423,12 @@ func FinishChecksum(seg []byte, hdr Hdr) error {
partial := binary.BigEndian.Uint16(seg[cs+co : cs+co+2]) partial := binary.BigEndian.Uint16(seg[cs+co : cs+co+2])
seg[cs+co] = 0 seg[cs+co] = 0
seg[cs+co+1] = 0 seg[cs+co+1] = 0
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], ^checksum.Checksum(seg[cs:], partial)) 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 return nil
} }
+284 -17
View File
@@ -226,13 +226,14 @@ func TestCorrectHdrLenChecksumBound(t *testing.T) {
// = 26) accepts. This case FAILS against the CsumStart+CsumStart regression. // = 26) accepts. This case FAILS against the CsumStart+CsumStart regression.
t.Run("valid-small-uso-accepted", func(t *testing.T) { t.Run("valid-small-uso-accepted", func(t *testing.T) {
pkt, _, csumStart := buildUDPv4Super(12) // total len 40 pkt, _, csumStart := buildUDPv4Super(12) // total len 40
hdr := Hdr{ hdr := NewHeader(
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4, unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
GSOSize: 6, // two 6-byte segments 0, /*hdrLen*/
CsumStart: csumStart, 6, /*gsoSize: two 6-byte segments*/
CsumOffset: 6, csumStart, /*csumStart*/
} 6, /*csumOffset*/
)
if err := CorrectHdrLen(pkt, &hdr); err != nil { if err := CorrectHdrLen(pkt, &hdr); err != nil {
t.Fatalf("CorrectHdrLen rejected a valid 40-byte USO superpacket: %v", err) t.Fatalf("CorrectHdrLen rejected a valid 40-byte USO superpacket: %v", err)
} }
@@ -247,13 +248,14 @@ func TestCorrectHdrLenChecksumBound(t *testing.T) {
t.Run("too-short-rejected", func(t *testing.T) { t.Run("too-short-rejected", func(t *testing.T) {
pkt := make([]byte, 25) pkt := make([]byte, 25)
pkt[0] = 0x45 // IPv4, IHL 5 pkt[0] = 0x45 // IPv4, IHL 5
hdr := Hdr{ hdr := NewHeader(
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4, unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
GSOSize: 6, 0, /*hdrLen*/
CsumStart: 20, 6, /*gsoSize*/
CsumOffset: 6, 20, /*csumStart*/
} 6, /*csumOffset*/
)
if err := CorrectHdrLen(pkt, &hdr); err == nil { if err := CorrectHdrLen(pkt, &hdr); err == nil {
t.Fatalf("CorrectHdrLen accepted a too-short (25-byte) packet") t.Fatalf("CorrectHdrLen accepted a too-short (25-byte) packet")
} }
@@ -303,9 +305,10 @@ func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 53 { if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 53 {
t.Errorf("seg %d: dport=%d want 53", i, dport) t.Errorf("seg %d: dport=%d want 53", i, dport)
} }
// UDP-GSO keeps the same IPv4 ID across every segment. // Software UDP GSO bumps the IPv4 ID per segment just like TSO
if id := binary.BigEndian.Uint16(seg[4:6]); id != 0x4242 { // (inet_gso_segment's fixed-ID case is TCP-only).
t.Errorf("seg %d: ip id=%#x want 0x4242", i, id) 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) segPayLen := len(seg) - int(hdrLen)
@@ -333,3 +336,267 @@ func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
}) })
} }
} }
// buildUDPv4Single constructs a single IPv4/UDP datagram with the checksum field pre-loaded with the folded
// pseudo-header sum, which is how a NEEDS_CSUM (CHECKSUM_PARTIAL) packet arrives from the tun.
func buildUDPv4Single(payload []byte) (pkt []byte, hdr Hdr) {
const ipLen, udpLen = 20, 8
pkt = make([]byte, ipLen+udpLen+len(payload))
pkt[0] = 0x45
binary.BigEndian.PutUint16(pkt[2:4], uint16(len(pkt)))
pkt[8] = 64
pkt[9] = unix.IPPROTO_UDP
copy(pkt[12:16], []byte{10, 0, 0, 1})
copy(pkt[16:20], []byte{10, 0, 0, 2})
binary.BigEndian.PutUint16(pkt[ipLen:ipLen+2], 12345)
binary.BigEndian.PutUint16(pkt[ipLen+2:ipLen+4], 53)
binary.BigEndian.PutUint16(pkt[ipLen+4:ipLen+6], uint16(udpLen+len(payload)))
copy(pkt[ipLen+udpLen:], payload)
pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, udpLen+len(payload))
binary.BigEndian.PutUint16(pkt[ipLen+udpChecksumOff:ipLen+udpChecksumOff+2], pseudo)
return pkt, NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, ipLen, udpChecksumOff)
}
// TestFinishChecksumUDPZeroStoresAllOnes pins RFC 768: a UDP checksum that computes to zero goes on the wire as
// 0xffff, because all-zero is the reserved "no checksum" encoding that IPv6 rejects outright.
func TestFinishChecksumUDPZeroStoresAllOnes(t *testing.T) {
var payload []byte
for i := 0; i < 0x10000; i++ {
p := []byte{byte(i >> 8), byte(i)}
pkt, hdr := buildUDPv4Single(p)
cs, co := int(hdr.CsumStart), int(hdr.CsumOffset)
partial := binary.BigEndian.Uint16(pkt[cs+co : cs+co+2])
pkt[cs+co], pkt[cs+co+1] = 0, 0
if ^checksum.Checksum(pkt[cs:], partial) == 0 {
payload = p
break
}
}
if payload == nil {
t.Fatal("no 2-byte payload produced a zero checksum")
}
pkt, hdr := buildUDPv4Single(payload)
if err := FinishChecksum(pkt, hdr); err != nil {
t.Fatalf("FinishChecksum: %v", err)
}
off := int(hdr.CsumStart) + int(hdr.CsumOffset)
if got := binary.BigEndian.Uint16(pkt[off : off+2]); got != 0xffff {
t.Fatalf("stored %#04x, want 0xffff for a UDP checksum that folds to zero", got)
}
}
// TestFinishChecksumTCPZeroPreserved is the other half of the protocol split: zero is a legal TCP checksum and must
// be stored as-is, so the UDP rewrite must not fire on a TCP csum_offset.
func TestFinishChecksumTCPZeroPreserved(t *testing.T) {
const cs, co = 20, tcpChecksumOff
// Seed the partial so the completed checksum lands on zero, the case the UDP path rewrites.
seg := make([]byte, cs+co+2)
for i := range seg[cs:] {
seg[cs+i] = byte(i * 7)
}
var partial uint16
for i := 0; i <= 0xffff; i++ {
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], uint16(i))
probe := append([]byte(nil), seg...)
probe[cs+co], probe[cs+co+1] = 0, 0
if ^checksum.Checksum(probe[cs:], uint16(i)) == 0 {
partial = uint16(i)
break
}
}
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], partial)
hdr := NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, cs, co)
if err := FinishChecksum(seg, hdr); err != nil {
t.Fatalf("FinishChecksum: %v", err)
}
if got := binary.BigEndian.Uint16(seg[cs+co : cs+co+2]); got != 0 {
t.Fatalf("stored %#04x, want 0x0000 (zero is a legal TCP checksum)", got)
}
}
// TestFinishChecksumUDPValidates confirms the ordinary path still produces a checksum a receiver accepts.
func TestFinishChecksumUDPValidates(t *testing.T) {
payload := []byte("the definitive tun offloads branch")
pkt, hdr := buildUDPv4Single(payload)
if err := FinishChecksum(pkt, hdr); err != nil {
t.Fatalf("FinishChecksum: %v", err)
}
pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, 8+len(payload))
if !verifyChecksum(pkt[hdr.CsumStart:], pseudo) {
t.Fatal("completed UDP checksum does not validate")
}
}
// TestCheckValidMasksGSOECN: GSO_ECN is a qualifier bit the kernel ORs
// into gso_type for TSO superpackets with CWR set. CheckValid must
// validate an ECN-qualified type as its base type — previously TCPV4|ECN
// fell into the default case and skipped the IP-version agreement check.
// The qualifier is TCP-only, so it must be rejected on UDP_L4.
func TestCheckValidMasksGSOECN(t *testing.T) {
v4pkt, _, _ := buildTCPv4Super(100)
v6pkt := make([]byte, len(v4pkt))
copy(v6pkt, v4pkt)
v6pkt[0] = 0x60 // claim IPv6
cases := []struct {
name string
pkt []byte
gsoType uint8
wantErr bool
}{
{"tcpv4-ecn-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, false},
{"tcpv4-ecn-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
{"tcpv6-ecn-v4-mismatch", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV6 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
{"udp-l4-ecn-rejected", v4pkt, unix.VIRTIO_NET_HDR_GSO_UDP_L4 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
{"tcpv4-plain-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, false},
{"tcpv4-plain-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, true},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
err := CheckValid(tc.pkt, NewHeader(0, tc.gsoType, 0, 100, 0, 0))
if tc.wantErr && err == nil {
t.Errorf("CheckValid(gsoType=%#x) = nil, want error", tc.gsoType)
}
if !tc.wantErr && err != nil {
t.Errorf("CheckValid(gsoType=%#x) = %v, want nil", tc.gsoType, err)
}
})
}
}
// TestCheckValidRejectsZeroGSOSize: a GSO-typed header with gso_size=0 must be
// rejected. It would produce a Packet whose GSOInfo.IsSuperpacket() is false,
// dodging both segmentation and FinishChecksum on its way downstream.
func TestCheckValidRejectsZeroGSOSize(t *testing.T) {
v4pkt, _, _ := buildTCPv4Super(100)
if err := CheckValid(v4pkt, NewHeader(0, unix.VIRTIO_NET_HDR_GSO_TCPV4, 0, 0, 0, 0)); err == nil {
t.Fatal("CheckValid accepted a GSO-typed header with gso_size=0")
}
}
// TestFoldComplementMatchesReference checks the segmenter's fold-and-invert
// against an independent RFC 1071 reference fold, hitting the carry edge
// cases (values whose first fold produces another carry).
func TestFoldComplementMatchesReference(t *testing.T) {
refFold := func(s uint64) uint16 {
for s>>16 != 0 {
s = s&0xffff + s>>16
}
return uint16(s)
}
cases := []uint32{
0, 1, 0xffff,
0x10000, // single carry
0x1fffe, // 0xffff + 0xffff: carry produces another 0xffff
0xffff0000, // high half only
0xfffeffff, // first fold yields another carry
0xffffffff, // worst case
}
for _, c := range cases {
if got, want := foldComplement(c), ^refFold(uint64(c)); got != want {
t.Errorf("foldComplement(%#x) = %#x, want %#x", c, got, want)
}
}
}
// referenceBaseIPv4HdrSum and referenceBaseTCPHdrSum are the straightforward
// implementations that baseIPv4HdrSum/baseTCPHdrSum replaced: copy the header
// into scratch, zero the fields the segment loop rewrites, sum. The production
// versions instead sum in place and subtract those fields via one's-complement
// arithmetic, which is faster but far less obvious — particularly for the TCP
// flags byte, which is only half of a 16-bit word. These references exist so
// that trade is checked rather than asserted.
func referenceBaseIPv4HdrSum(pkt []byte, ihl int) uint32 {
var ipTmp [ipv4HeaderMaxLen]byte
copy(ipTmp[:ihl], pkt[:ihl])
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0
return uint32(checksum.Checksum(ipTmp[:ihl], 0))
}
func referenceBaseTCPHdrSum(pkt []byte, csumStart, headerLen int) uint32 {
tcpLen := headerLen - csumStart
var tmp [tcpHeaderMaxLen]byte
copy(tmp[:tcpLen], pkt[csumStart:headerLen])
tmp[tcpSeqOff], tmp[tcpSeqOff+1], tmp[tcpSeqOff+2], tmp[tcpSeqOff+3] = 0, 0, 0, 0
tmp[tcpFlagsOff] = 0
tmp[tcpChecksumOff], tmp[tcpChecksumOff+1] = 0, 0
return uint32(checksum.Checksum(tmp[:tcpLen], 0))
}
// randSeed is a tiny deterministic PRNG so this test needs no imports beyond
// what the file already has and reproduces identically on every run.
func randByte(state *uint32) byte {
*state = *state*1664525 + 1013904223
return byte(*state >> 24)
}
func TestBaseSumsMatchZeroingReference(t *testing.T) {
state := uint32(12345)
t.Run("ipv4", func(t *testing.T) {
for ihl := ipv4HeaderMinLen; ihl <= ipv4HeaderMaxLen; ihl += 4 {
for iter := 0; iter < 5000; iter++ {
pkt := make([]byte, ihl)
for i := range pkt {
pkt[i] = randByte(&state)
}
pkt[0] = byte(0x40 | (ihl / 4))
want := referenceBaseIPv4HdrSum(pkt, ihl)
got, err := baseIPv4HdrSum(pkt, ihl)
if err != nil {
t.Fatalf("ihl=%d: %v", ihl, err)
}
// Compare the value that reaches the wire: the raw partial
// sums may legally differ by one's-complement -0 vs +0.
for _, tl := range []uint32{20, 1500, 65535} {
for _, id := range []uint32{0, 0x4242, 0xffff} {
if a, b := foldComplement(want+tl+id), foldComplement(got+tl+id); a != b {
t.Fatalf("ihl=%d tl=%d id=%d: %#04x != %#04x", ihl, tl, id, a, b)
}
}
}
}
}
})
t.Run("tcp", func(t *testing.T) {
const csumStart = 20
for dataOff := 5; dataOff <= 15; dataOff++ {
tcpLen := dataOff * 4
headerLen := csumStart + tcpLen
for iter := 0; iter < 5000; iter++ {
pkt := make([]byte, headerLen+64)
for i := range pkt {
pkt[i] = randByte(&state)
}
pkt[0] = 0x45
pkt[csumStart+tcpDataOffOff] = byte(dataOff << 4)
want := referenceBaseTCPHdrSum(pkt, csumStart, headerLen)
got := baseTCPHdrSum(pkt, csumStart, headerLen)
for _, seq := range []uint32{0, 1, 0x4242_4242, 0xffff_ffff} {
for _, fl := range []uint32{0x00, 0x10, 0x18, 0x19, 0xff} {
for _, l4 := range []uint32{20, 1460, 65535} {
a := foldComplement(want + seq + fl + l4)
b := foldComplement(got + seq + fl + l4)
if a != b {
t.Fatalf("dataOff=%d seq=%#x fl=%#x l4=%d: %#04x != %#04x",
dataOff, seq, fl, l4, a, b)
}
}
}
}
}
}
})
}
+3 -1
View File
@@ -550,7 +550,9 @@ func (t *tun) Read(to []byte) (int, error) {
return n - 4, nil return n - 4, nil
} }
// Write pushes one IP packet onto the utun device. Only valid for single threaded use. // 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) { func (t *tun) Write(from []byte) (int, error) {
if len(from) == 0 { if len(from) == 0 {
return 0, syscall.EIO return 0, syscall.EIO
+33 -82
View File
@@ -35,22 +35,7 @@ type tun struct {
deviceIndex int deviceIndex int
ioctlFd uintptr ioctlFd uintptr
vnetHdr bool vnetHdr bool
// offloadFlags is the exact TUN_F_* offload mask newTun negotiated with
// the kernel: usoOffloadFlags when USO was accepted, tsoOffloadFlags on
// the TSO-only fallback, or 0 when vnetHdr is off. TUNSETOFFLOAD is
// device-wide (drivers/net/tun.c set_offload updates tun->set_features
// for the whole netdev), so addQueue must replay this exact
// mask on every added queue — issuing a narrower mask there would
// silently downgrade offloads (e.g. disable USO) for all queues while
// they still advertise the stale capability.
offloadFlags uint offloadFlags uint
// routeFeatureECN, when true, sets RTAX_FEATURE_ECN on every route we
// install for the tun. The kernel then actively negotiates ECN for
// connections destined to those prefixes (equivalent to `ip route
// change ... features ecn`) regardless of net.ipv4.tcp_ecn, so flows
// across the nebula mesh use ECN even when the host default is the
// passive setting (=2). Disable via tunnels.ecn=false.
routeFeatureECN bool
Routes atomic.Pointer[[]Route] Routes atomic.Pointer[[]Route]
routeTree atomic.Pointer[bart.Table[routing.Gateways]] routeTree atomic.Pointer[bart.Table[routing.Gateways]]
@@ -91,14 +76,7 @@ type ifreqQLEN struct {
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) { func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
// We don't know what flags the caller opened this fd with and can't turn // 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. // on IFF_VNET_HDR after TUNSETIFF, so skip offload on inherited fds.
t, err := newTunGeneric(c, l, deviceFd, false, 0, vpnNetworks) return newTunGeneric(c, l, deviceFd, false, 0, vpnNetworks, "tun0")
if err != nil {
return nil, err
}
t.Device = "tun0"
return t, nil
} }
// openTunDev opens /dev/net/tun, creating the device node first if it's // openTunDev opens /dev/net/tun, creating the device node first if it's
@@ -124,8 +102,7 @@ func openTunDev() (int, error) {
return fd, nil return fd, nil
} }
// tunSetIff runs TUNSETIFF with the given flags and returns the kernel-chosen // tunSetIff runs TUNSETIFF with the given flags and returns the kernel-chosen device name on success.
// device name on success.
func tunSetIff(fd int, name string, flags uint16) (string, error) { func tunSetIff(fd int, name string, flags uint16) (string, error) {
var req ifReq var req ifReq
req.Flags = flags req.Flags = flags
@@ -136,57 +113,45 @@ func tunSetIff(fd int, name string, flags uint16) (string, error) {
return strings.Trim(string(req.Name[:]), "\x00"), nil return strings.Trim(string(req.Name[:]), "\x00"), nil
} }
// tsoOffloadFlags are the TUN_F_* bits we ask the kernel to enable when a // tsoOffloadFlags are the TUN_F_* bits we ask the kernel to enable when a TSO-capable TUN is available.
// TSO-capable TUN is available. CSUM is required as a prerequisite for TSO.
// TSO_ECN tells the kernel we propagate ECN correctly through coalesce and
// segmentation, so it can deliver superpackets whose seed has CWR/ECE set
// or whose IP-level codepoint is CE.
const tsoOffloadFlags = unix.TUN_F_CSUM | unix.TUN_F_TSO4 | unix.TUN_F_TSO6 | unix.TUN_F_TSO_ECN const tsoOffloadFlags = unix.TUN_F_CSUM | unix.TUN_F_TSO4 | unix.TUN_F_TSO6 | unix.TUN_F_TSO_ECN
// usoOffloadFlags adds UDP Segmentation Offload to tsoOffloadFlags. Requires // usoAndTSOOffloadFlags adds UDP Segmentation Offload to tsoOffloadFlags.
// Linux 6.2; older kernels reject it and we fall back to TCP-only TSO via // Requires Linux >= 6.2; older kernels reject it and we fall back to TCP-only TSO
// tsoOffloadFlags. const usoAndTSOOffloadFlags = tsoOffloadFlags | unix.TUN_F_USO4 | unix.TUN_F_USO6
const usoOffloadFlags = tsoOffloadFlags | unix.TUN_F_USO4 | unix.TUN_F_USO6
// offloadUSOEnabled reports whether the negotiated offload mask includes UDP
// Segmentation Offload. It is the single source of truth for the usoEnabled
// capability surfaced by each queue, so the mask stored on the tun and the USO
// bit reported to coalescers can never drift apart.
func offloadUSOEnabled(offloadFlags uint) bool { func offloadUSOEnabled(offloadFlags uint) bool {
return offloadFlags&(unix.TUN_F_USO4|unix.TUN_F_USO6) != 0 return offloadFlags&(unix.TUN_F_USO4|unix.TUN_F_USO6) != 0
} }
func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) { func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) {
baseFlags := uint16(unix.IFF_TUN | unix.IFF_NO_PI) // 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 { if multiqueue {
baseFlags |= unix.IFF_MULTI_QUEUE baseFlags |= unix.IFF_MULTI_QUEUE
} }
nameStr := c.GetString("tun.dev", "") nameStr := c.GetString("tun.dev", "")
// First try to enable IFF_VNET_HDR via TUNSETIFF and negotiate TUN_F_*
// offloads via TUNSETOFFLOAD so we can receive TSO/USO superpackets.
// We try TSO+USO first, fall back to TSO-only on kernels without USO
// (Linux < 6.2), and finally give up on virtio headers entirely and
// reopen as a plain TUN if neither offload mask is accepted.
fd, err := openTunDev() fd, err := openTunDev()
if err != nil { if err != nil {
return nil, err return nil, err
} }
vnetHdr := true vnetHdr := true
// offloadFlags is the exact TUN_F_* mask the kernel accepted. We remember
// it (rather than a plain bool) so addQueue can replay the // First try to enable IFF_VNET_HDR via TUNSETIFF and negotiate TUN_F_* offloads
// identical device-wide mask on added queues instead of downgrading them. // 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
var offloadFlags uint var offloadFlags uint
name, err := tunSetIff(fd, nameStr, baseFlags|unix.IFF_VNET_HDR) name, err := tunSetIff(fd, nameStr, baseFlags|unix.IFF_VNET_HDR)
if err != nil { if err != nil {
_ = unix.Close(fd) _ = unix.Close(fd)
vnetHdr = false vnetHdr = false
} else { } else {
// Try TSO+USO first. On kernels without USO support (Linux < 6.2) if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(usoAndTSOOffloadFlags)); err == nil {
// the ioctl returns EINVAL; fall back to the TCP-only mask before offloadFlags = usoAndTSOOffloadFlags
// giving up on VNET_HDR entirely.
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(usoOffloadFlags)); err == nil {
offloadFlags = usoOffloadFlags
} else if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(tsoOffloadFlags)); err == nil { } else if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(tsoOffloadFlags)); err == nil {
offloadFlags = tsoOffloadFlags offloadFlags = tsoOffloadFlags
} else { } else {
@@ -212,7 +177,7 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
l.Info("TUN offload enabled", "tso", true, "uso", offloadUSOEnabled(offloadFlags)) l.Info("TUN offload enabled", "tso", true, "uso", offloadUSOEnabled(offloadFlags))
} }
t, err := newTunGeneric(c, l, fd, vnetHdr, offloadFlags, vpnNetworks) t, err := newTunGeneric(c, l, fd, vnetHdr, offloadFlags, vpnNetworks, name)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -222,16 +187,14 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
return t, nil return t, nil
} }
// newTunGeneric does all the stuff common to different tun initialization // newTunGeneric does all the stuff common to different tun initialization paths.
// paths. It will close your files on error. offloadFlags is the TUN_F_* mask // It will close your files on error.
// newTun negotiated (0 when vnetHdr is off); the queues' USO capability is // offloadFlags is the TUN_F_* mask newTun negotiated (ignored when vnetHdr is false)
// derived from it so it can never disagree with the mask we replay on added func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFlags uint, vpnNetworks []netip.Prefix, name string) (*tun, error) {
// multiqueue readers.
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFlags uint, vpnNetworks []netip.Prefix) (*tun, error) {
var qs tio.QueueSet var qs tio.QueueSet
var err error var err error
if vnetHdr { if vnetHdr {
qs, err = tio.NewOffloadQueueSet(offloadUSOEnabled(offloadFlags)) qs, err = tio.NewOffloadQueueSet(offloadUSOEnabled(offloadFlags), l)
} else { } else {
qs, err = tio.NewPollQueueSet() qs, err = tio.NewPollQueueSet()
} }
@@ -242,11 +205,15 @@ func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFla
} }
err = qs.Add(fd) err = qs.Add(fd)
if err != nil { 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) _ = unix.Close(fd)
_ = qs.Close()
return nil, err return nil, err
} }
t := &tun{ t := &tun{
Device: name,
readers: qs, readers: qs,
closeLock: sync.Mutex{}, closeLock: sync.Mutex{},
vnetHdr: vnetHdr, vnetHdr: vnetHdr,
@@ -255,7 +222,6 @@ func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr bool, offloadFla
TXQueueLen: c.GetInt("tun.tx_queue", 500), TXQueueLen: c.GetInt("tun.tx_queue", 500),
useSystemRoutes: c.GetBool("tun.use_system_route_table", false), useSystemRoutes: c.GetBool("tun.use_system_route_table", false),
useSystemRoutesBufferSize: c.GetInt("tun.use_system_route_table_buffer_size", 0), useSystemRoutesBufferSize: c.GetInt("tun.use_system_route_table_buffer_size", 0),
routeFeatureECN: c.GetBool("tunnels.ecn", true),
routesFromSystem: map[netip.Prefix]routing.Gateways{}, routesFromSystem: map[netip.Prefix]routing.Gateways{},
l: l, l: l,
} }
@@ -349,9 +315,7 @@ func (t *tun) reload(c *config.C, initial bool) error {
return nil return nil
} }
// Queues opens additional kernel multiqueue fds until the device has n // Queues opens additional kernel multiqueue fds until the device has n queues, then returns them all.
// queues, then returns them all. The first queue was opened by newTun; each
// extra fd replays the negotiated offload state (see addQueue).
func (t *tun) Queues(n int) ([]tio.Queue, error) { func (t *tun) Queues(n int) ([]tio.Queue, error) {
for len(t.readers.Queues()) < n { for len(t.readers.Queues()) < n {
if err := t.addQueue(); err != nil { if err := t.addQueue(); err != nil {
@@ -361,8 +325,7 @@ func (t *tun) Queues(n int) ([]tio.Queue, error) {
return t.readers.Queues(), nil return t.readers.Queues(), nil
} }
// addQueue opens one more IFF_MULTI_QUEUE fd on the device and adds it to // addQueue opens one more IFF_MULTI_QUEUE fd on the device and adds it to the queue set.
// the queue set.
func (t *tun) addQueue() error { func (t *tun) addQueue() error {
t.closeLock.Lock() t.closeLock.Lock()
defer t.closeLock.Unlock() defer t.closeLock.Unlock()
@@ -382,10 +345,6 @@ func (t *tun) addQueue() error {
} }
if t.vnetHdr { if t.vnetHdr {
// Replay the exact mask newTun negotiated. TUNSETOFFLOAD is
// device-wide, so issuing the TSO-only mask here would disable USO
// for every queue (including queue 0) on kernels where newTun
// successfully enabled it, while the queues keep advertising USO.
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(t.offloadFlags)); err != nil { if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(t.offloadFlags)); err != nil {
_ = unix.Close(fd) _ = unix.Close(fd)
return fmt.Errorf("failed to enable offload on multiqueue tun fd: %w", err) return fmt.Errorf("failed to enable offload on multiqueue tun fd: %w", err)
@@ -566,18 +525,13 @@ func (t *tun) setDefaultRoute(cidr netip.Prefix) error {
Table: unix.RT_TABLE_MAIN, Table: unix.RT_TABLE_MAIN,
Type: unix.RTN_UNICAST, Type: unix.RTN_UNICAST,
} }
// Match the metric the kernel uses for its auto-installed connected // Match the metric the kernel uses for its auto-installed connected route,
// route, so RouteReplace overwrites it in place instead of adding a // so RouteReplace overwrites it in place instead of adding a second route at a worse metric.
// second route at a worse metric. IPv6 connected routes are installed // IPv6 connected routes are installed at metric 256 (IP6_RT_PRIO_KERN); IPv4 uses 0.
// at metric 256 (IP6_RT_PRIO_KERN); IPv4 uses 0. Without this, the // Without this, the kernel route wins lookups and our MTU / AdvMSS / Features never apply on v6.
// kernel route wins lookups and our MTU / AdvMSS / Features never
// apply on v6.
if cidr.Addr().Is6() { if cidr.Addr().Is6() {
nr.Priority = 256 nr.Priority = 256
} }
if t.routeFeatureECN {
nr.Features |= unix.RTAX_FEATURE_ECN
}
err := netlink.RouteReplace(&nr) err := netlink.RouteReplace(&nr)
if err != nil { if err != nil {
t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr) t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr)
@@ -627,9 +581,6 @@ func (t *tun) addRoutes(logErrors bool) error {
if r.Metric > 0 { if r.Metric > 0 {
nr.Priority = r.Metric nr.Priority = r.Metric
} }
if t.routeFeatureECN {
nr.Features |= unix.RTAX_FEATURE_ECN
}
err := netlink.RouteReplace(&nr) err := netlink.RouteReplace(&nr)
if err != nil { if err != nil {
+9 -9
View File
@@ -39,14 +39,14 @@ func TestTunAdvMSS(t *testing.T) {
// capability: it is derived from the negotiated offload mask, so the mask // capability: it is derived from the negotiated offload mask, so the mask
// stored on the tun and the capability reported to coalescers cannot drift. // stored on the tun and the capability reported to coalescers cannot drift.
func TestOffloadUSOEnabled(t *testing.T) { func TestOffloadUSOEnabled(t *testing.T) {
// usoOffloadFlags must be a strict superset of tsoOffloadFlags. Otherwise // usoAndTSOOffloadFlags must be a strict superset of tsoOffloadFlags. Otherwise
// the TSO-only fallback (and the historic hardcoded-mask bug in // the TSO-only fallback (and the historic hardcoded-mask bug in
// addQueue) would not actually be a downgrade. // addQueue) would not actually be a downgrade.
if usoOffloadFlags&tsoOffloadFlags != tsoOffloadFlags { if usoAndTSOOffloadFlags&tsoOffloadFlags != tsoOffloadFlags {
t.Fatalf("usoOffloadFlags (%#x) is not a superset of tsoOffloadFlags (%#x)", usoOffloadFlags, tsoOffloadFlags) t.Fatalf("usoAndTSOOffloadFlags (%#x) is not a superset of tsoOffloadFlags (%#x)", usoAndTSOOffloadFlags, tsoOffloadFlags)
} }
if usoOffloadFlags == tsoOffloadFlags { if usoAndTSOOffloadFlags == tsoOffloadFlags {
t.Fatal("usoOffloadFlags must add bits beyond tsoOffloadFlags") t.Fatal("usoAndTSOOffloadFlags must add bits beyond tsoOffloadFlags")
} }
cases := []struct { cases := []struct {
@@ -54,7 +54,7 @@ func TestOffloadUSOEnabled(t *testing.T) {
offloadFlags uint offloadFlags uint
wantUSO bool wantUSO bool
}{ }{
{"uso-negotiated", usoOffloadFlags, true}, {"uso-negotiated", usoAndTSOOffloadFlags, true},
{"tso-fallback", tsoOffloadFlags, false}, {"tso-fallback", tsoOffloadFlags, false},
{"no-vnet-hdr", 0, false}, {"no-vnet-hdr", 0, false},
} }
@@ -78,12 +78,12 @@ func TestOffloadUSOEnabled(t *testing.T) {
// TUNSETOFFLOAD argument is read from. // TUNSETOFFLOAD argument is read from.
func TestAddQueueReplaysNegotiatedMask(t *testing.T) { func TestAddQueueReplaysNegotiatedMask(t *testing.T) {
t.Run("uso-negotiated", func(t *testing.T) { t.Run("uso-negotiated", func(t *testing.T) {
tn := &tun{vnetHdr: true, offloadFlags: usoOffloadFlags} tn := &tun{vnetHdr: true, offloadFlags: usoAndTSOOffloadFlags}
// The ioctl argument in addQueue is uintptr(t.offloadFlags); // The ioctl argument in addQueue is uintptr(t.offloadFlags);
// it must equal the negotiated USO mask, and must NOT be the TSO-only // it must equal the negotiated USO mask, and must NOT be the TSO-only
// mask (the original bug). // mask (the original bug).
if tn.offloadFlags != usoOffloadFlags { if tn.offloadFlags != usoAndTSOOffloadFlags {
t.Fatalf("offloadFlags = %#x, want %#x", tn.offloadFlags, usoOffloadFlags) t.Fatalf("offloadFlags = %#x, want %#x", tn.offloadFlags, usoAndTSOOffloadFlags)
} }
if tn.offloadFlags == tsoOffloadFlags { if tn.offloadFlags == tsoOffloadFlags {
t.Fatal("added queue would downgrade USO: offloadFlags must not be the TSO-only mask when USO was negotiated") t.Fatal("added queue would downgrade USO: offloadFlags must not be the TSO-only mask when USO was negotiated")
+6 -4
View File
@@ -10,11 +10,13 @@ import (
_ "net/http/pprof" // registers pprof handlers on http.DefaultServeMux _ "net/http/pprof" // registers pprof handlers on http.DefaultServeMux
) )
// startPprofServer serves net/http/pprof on :6060 for the life of ctx. It is // startPprofServer serves net/http/pprof on localhost:6060 for the life of
// only compiled into debug builds (`-tags debug`, `make debug`), so a debug // ctx. It is only compiled into debug builds (`-tags debug`, `make debug`),
// build announces itself with the Info line below. // so a debug build announces itself with the Info line below. Loopback only:
// a wildcard bind would expose profiles (peer addresses, config-derived
// state) to anything that can reach the host, the overlay included.
func startPprofServer(ctx context.Context, l *slog.Logger) { func startPprofServer(ctx context.Context, l *slog.Logger) {
server := &http.Server{Addr: ":6060", Handler: nil} server := &http.Server{Addr: "localhost:6060", Handler: nil}
l.Info("Starting pprof debug server (debug build)", "addr", server.Addr) l.Info("Starting pprof debug server (debug build)", "addr", server.Addr)
go func() { go func() {
+1 -1
View File
@@ -161,7 +161,7 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak
switch existingRelay.State { switch existingRelay.State {
case Established: case Established:
hl.Log(context.Background(), level, "Send handshake via relay", "relay", relay.String()) hl.Log(context.Background(), level, "Send handshake via relay", "relay", relay.String())
f.SendVia(relayHostInfo, existingRelay, stage0, make([]byte, 12), make([]byte, mtu), false) f.SendVia(relayHostInfo, existingRelay, stage0, make([]byte, 12), make([]byte, mtu), false, 0)
case Disestablished: case Disestablished:
// Mark this relay as 'requested' // Mark this relay as 'requested'
relayHostInfo.relayState.UpdateRelayForByIpState(vpnIp, Requested) relayHostInfo.relayState.UpdateRelayForByIpState(vpnIp, Requested)
+13 -28
View File
@@ -14,43 +14,28 @@ const MTU = 9001
// only costs additional sendmmsg chunks within a single WriteBatch call. // only costs additional sendmmsg chunks within a single WriteBatch call.
const MaxWriteBatch = 128 const MaxWriteBatch = 128
// RxMeta carries per-packet metadata extracted from the RX path (ancillary
// data, kernel offload state, etc.) and passed to EncReader callbacks.
// Backends that do not produce a particular signal leave its zero value.
//
// OuterECN is the 2-bit IP-level ECN codepoint stamped on the carrier
// datagram (extracted from IP_TOS / IPV6_TCLASS cmsg on Linux). Zero
// means Not-ECT, which is also the value backends without ECN RX support
// supply on every packet.
type RxMeta struct {
OuterECN byte
}
type EncReader func( type EncReader func(
addr netip.AddrPort, addr netip.AddrPort,
payload []byte, payload []byte,
meta RxMeta,
) )
type Conn interface { type Conn interface {
Rebind() error Rebind() error
LocalAddr() (netip.AddrPort, error) LocalAddr() (netip.AddrPort, error)
// ListenOut invokes r for each received packet. On batch-capable // ListenOut invokes r for each received packet.
// backends (recvmmsg), flush is called after each batch is fully // On batch-capable backends (recvmmsg), flush is called after each batch is fully delivered.
// delivered — callers use it to flush per-batch accumulators such as // Callers use it to flush per-batch accumulators such as TUN write coalescers.
// TUN write coalescers. Single-packet backends call flush after each // Single-packet backends call flush after each packet. flush must not be nil.
// packet. flush must not be nil.
ListenOut(r EncReader, flush func()) error ListenOut(r EncReader, flush func()) error
WriteTo(b []byte, addr netip.AddrPort) error WriteTo(b []byte, addr netip.AddrPort) error
// WriteBatch sends a contiguous batch of packets, each with its own // WriteBatch sends a contiguous batch of packets, each with its own
// destination. bufs and addrs must have the same length. outerECNs may // destination. bufs and addrs must have the same length. Linux uses
// be nil (treated as all-zero / Not-ECT); when non-nil it must have the // sendmmsg(2) for a single syscall.
// same length as bufs, and outerECNs[i] is the 2-bit IP-level ECN //
// codepoint to set on packet i's outer header. Linux uses sendmmsg(2) // Returns the number of packets successfully written. A destination the kernel rejects costs only
// for a single syscall and attaches the value as IP_TOS / IPV6_TCLASS // its own packet, so a short count means some peers were undeliverable, not that the batch failed.
// cmsg; other backends ignore it. Returns on the first error; callers // Not safe for concurrent use on the same Conn.
// may observe a partial send if some packets went out before the error. WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error)
WriteBatch(bufs [][]byte, addrs []netip.AddrPort, outerECNs []byte) error
ReloadConfig(c *config.C) ReloadConfig(c *config.C)
SupportsMultipleReaders() bool SupportsMultipleReaders() bool
Close() error Close() error
@@ -73,8 +58,8 @@ func (NoopConn) SupportsMultipleReaders() bool {
func (NoopConn) WriteTo(_ []byte, _ netip.AddrPort) error { func (NoopConn) WriteTo(_ []byte, _ netip.AddrPort) error {
return nil return nil
} }
func (NoopConn) WriteBatch(_ [][]byte, _ []netip.AddrPort, _ []byte) error { func (NoopConn) WriteBatch(bufs [][]byte, _ []netip.AddrPort) (int, error) {
return nil return len(bufs), nil
} }
func (NoopConn) ReloadConfig(_ *config.C) { func (NoopConn) ReloadConfig(_ *config.C) {
return return
+61
View File
@@ -0,0 +1,61 @@
package udp
import (
"context"
"log/slog"
"github.com/slackhq/nebula/config"
)
// NetworkChangeMonitor rebinds the udp listener when the local network moves out from under it.
//
// Detection lives here in the udp package, next to the socket it concerns and the platform matrix that already knows
// which sockets go stale. What to do about a change — updating the lighthouse, requerying tunnels — is not the udp
// package's business, so Start takes the reaction as a plain function. Passing it at Start rather than holding it
// keeps this package from referencing whatever owns the rebind.
//
// On platforms whose sockets do not go stale, watchNetworkChanges hands back a nil channel and Start returns.
type NetworkChangeMonitor struct {
l *slog.Logger
ctx context.Context
enabled bool
}
// NewNetworkChangeMonitor builds a monitor for local network changes. The returned monitor is always usable: Start
// is safe to call unconditionally, it no-ops when disabled or on a platform that does not need it.
func NewNetworkChangeMonitor(ctx context.Context, l *slog.Logger, c *config.C) *NetworkChangeMonitor {
return &NetworkChangeMonitor{
l: l,
ctx: ctx,
enabled: c.GetBool("listen.rebind_on_network_change", true),
}
}
// Start watches for network changes until the context is cancelled, calling rebind once per settled change. It
// blocks, so callers run it in a goroutine, and it no-ops when disabled, unsupported, or with nothing to rebind.
func (m *NetworkChangeMonitor) Start(rebind func()) {
if !m.enabled || rebind == nil || m.ctx.Err() != nil {
return
}
changes, err := watchNetworkChanges(m.ctx, m.l)
if err != nil {
// Not fatal. Everything else still works, we just won't notice a network change on our own.
m.l.Error("Failed to watch for network changes, will not rebind the udp listener when the network moves",
"error", err,
)
return
}
if changes == nil {
// This platform's sockets don't go stale, so there is nothing to watch for.
return
}
m.l.Info("Watching for network changes to rebind the udp listener")
for range changes {
m.l.Info("Local network changed, rebinding the udp listener")
rebind()
}
}
+164
View File
@@ -0,0 +1,164 @@
//go:build darwin && !ios && !e2e_testing
// +build darwin,!ios,!e2e_testing
package udp
import (
"context"
"encoding/binary"
"errors"
"log/slog"
"os"
"time"
"golang.org/x/sys/unix"
)
const (
// netChangeSettleWindow is how long we keep swallowing routing messages after the first interesting one. A
// single network change is never a single message, it is a burst: the link drops, addresses go away, new ones
// arrive, routes get rewritten. Reporting part way through that just means reporting again.
netChangeSettleWindow = time.Second
// netChangeReadBuffer is sized well past any rt_msghdr plus its addresses. A short read would be discarded by
// the kernel, so being generous here is how we avoid missing a message.
netChangeReadBuffer = 4096
)
// watchNetworkChanges reports when the local network moves out from under us, so the listener can be rebound.
//
// Darwin scopes a udp socket to whatever interface it came up on. Move between networks and we keep sending out an
// interface that no longer has a route, which surfaces as an instant "no route to host" with no packet ever leaving
// the box. Rebind clears that, but only if something notices the change and calls it. iOS has always been told by
// the host app off NWPathMonitor. This is the equivalent for everything else that runs on darwin.
//
// The returned channel is buffered and coalescing: a send is dropped if one is already pending, since both mean the
// same thing to a reader. It is closed when ctx is cancelled or the routing socket fails, so a caller can simply
// range over it. Platforms whose sockets do not need rebinding return a nil channel and no error.
func watchNetworkChanges(ctx context.Context, l *slog.Logger) (<-chan struct{}, error) {
sock, err := openRouteSocket()
if err != nil {
return nil, err
}
changes := make(chan struct{}, 1)
go func() {
defer close(changes)
defer func() { _ = sock.Close() }()
// Closing the socket is what unblocks the read in watchRouteSocket, so this turns cancellation into a
// close. It is scoped to this call so it cannot outlive the watch it belongs to.
done := make(chan struct{})
defer close(done)
go func() {
select {
case <-ctx.Done():
_ = sock.Close()
case <-done:
}
}()
watchRouteSocket(l, sock, changes)
}()
return changes, nil
}
// watchRouteSocket blocks reading the routing socket, reporting once per settled burst of changes. It returns when
// the socket is closed, which is how cancellation gets us out of here.
func watchRouteSocket(l *slog.Logger, sock *os.File, changes chan<- struct{}) {
buf := make([]byte, netChangeReadBuffer)
for {
n, err := sock.Read(buf)
if err != nil {
logRouteSocketError(l, err)
return
}
if !isNetworkChange(buf[:n]) {
continue
}
// Swallow the rest of the burst. The deadline is absolute and not extended by what arrives, so this always
// ends after the settle window no matter how chatty the socket is. Changes that land after the window
// simply produce another report, which is the correct outcome anyway.
deadline := time.Now().Add(netChangeSettleWindow)
for {
if err = sock.SetReadDeadline(deadline); err != nil {
logRouteSocketError(l, err)
return
}
if _, err = sock.Read(buf); err != nil {
if os.IsTimeout(err) {
break
}
logRouteSocketError(l, err)
return
}
}
if err = sock.SetReadDeadline(time.Time{}); err != nil {
logRouteSocketError(l, err)
return
}
select {
case changes <- struct{}{}:
default:
// One already pending, and a second "the network moved" tells the reader nothing new.
}
}
}
// logRouteSocketError reports a routing socket failure unless it is just us shutting the socket down.
func logRouteSocketError(l *slog.Logger, err error) {
if errors.Is(err, os.ErrClosed) {
return
}
l.Error("Error reading the routing socket, will no longer notice local network changes", "error", err)
}
// openRouteSocket returns the routing socket as a non blocking os.File. Going through os.File puts reads on the go
// poller, which buys us both a working read deadline and a Close that unblocks a read in progress.
func openRouteSocket() (*os.File, error) {
fd, err := unix.Socket(unix.AF_ROUTE, unix.SOCK_RAW, unix.AF_UNSPEC)
if err != nil {
return nil, err
}
if err = unix.SetNonblock(fd, true); err != nil {
_ = unix.Close(fd)
return nil, err
}
return os.NewFile(uintptr(fd), "route"), nil
}
// isNetworkChange reports whether a routing message means our local addressing may have moved out from under us.
//
// We read the header instead of parsing the message because the type is the only part we need, and a full parse can
// fail on shapes we don't care about, which would turn "a message I can't parse" into "a change I missed".
// rt_msghdr, if_msghdr and ifa_msghdr all begin with the same three fields, so this is the same for every type.
func isNetworkChange(msg []byte) bool {
if len(msg) < 4 {
return false
}
// u_short msglen, u_char version, u_char type
if int(binary.NativeEndian.Uint16(msg[0:2])) > len(msg) || msg[2] != unix.RTM_VERSION {
return false
}
switch msg[3] {
case unix.RTM_NEWADDR, unix.RTM_DELADDR, unix.RTM_IFINFO:
// An address arrived or left, or a link changed state. Anything else on this socket is either a route
// churning underneath us, which a rebind doesn't help with, or unrelated traffic.
return true
default:
return false
}
}
+244
View File
@@ -0,0 +1,244 @@
//go:build darwin && !ios && !e2e_testing
// +build darwin,!ios,!e2e_testing
package udp
import (
"context"
"encoding/binary"
"os"
"testing"
"time"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/test"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/goleak"
"golang.org/x/sys/unix"
)
// routeMsg builds the first four bytes of a routing message, which is all isNetworkChange reads.
func routeMsg(msgType uint8, extra int) []byte {
msg := make([]byte, 4+extra)
binary.NativeEndian.PutUint16(msg[0:2], uint16(len(msg)))
msg[2] = unix.RTM_VERSION
msg[3] = msgType
return msg
}
func TestIsNetworkChange(t *testing.T) {
// The three that mean our addressing may have moved
assert.True(t, isNetworkChange(routeMsg(unix.RTM_NEWADDR, 0)))
assert.True(t, isNetworkChange(routeMsg(unix.RTM_DELADDR, 0)))
assert.True(t, isNetworkChange(routeMsg(unix.RTM_IFINFO, 0)))
// Route churn is not something a rebind helps with
assert.False(t, isNetworkChange(routeMsg(unix.RTM_ADD, 0)))
assert.False(t, isNetworkChange(routeMsg(unix.RTM_DELETE, 0)))
assert.False(t, isNetworkChange(routeMsg(unix.RTM_GET, 0)))
// Garbage must not be mistaken for a change
assert.False(t, isNetworkChange(nil), "empty")
assert.False(t, isNetworkChange([]byte{0, 0, 0}), "short header")
wrongVersion := routeMsg(unix.RTM_NEWADDR, 0)
wrongVersion[2] = unix.RTM_VERSION + 1
assert.False(t, isNetworkChange(wrongVersion), "wrong rtm_version")
lying := routeMsg(unix.RTM_NEWADDR, 0)
binary.NativeEndian.PutUint16(lying[0:2], 512)
assert.False(t, isNetworkChange(lying), "msglen longer than what we read")
}
// socketPair returns a connected pair of datagram sockets, the first wrapped the same way the routing socket is. It
// stands in for the kernel so the watch loop can be driven with synthetic messages.
func socketPair(t *testing.T) (*os.File, int) {
t.Helper()
fds, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_DGRAM, 0)
require.NoError(t, err)
require.NoError(t, unix.SetNonblock(fds[0], true))
f := os.NewFile(uintptr(fds[0]), "route")
t.Cleanup(func() {
_ = f.Close()
_ = unix.Close(fds[1])
})
return f, fds[1]
}
func TestWatchRouteSocketCoalescesABurst(t *testing.T) {
sock, kernel := socketPair(t)
changes := make(chan struct{}, 1)
done := make(chan struct{})
go func() {
watchRouteSocket(test.NewLogger(), sock, changes)
close(done)
}()
// One network change is a burst of messages. All of these land inside the settle window, so they must produce
// exactly one report rather than one apiece.
for range 5 {
_, err := unix.Write(kernel, routeMsg(unix.RTM_NEWADDR, 8))
require.NoError(t, err)
}
// Uninteresting messages in the middle of a burst must not add a report of their own either.
_, err := unix.Write(kernel, routeMsg(unix.RTM_ADD, 8))
require.NoError(t, err)
select {
case <-changes:
case <-time.After(netChangeSettleWindow * 4):
t.Fatal("a burst should have reported a change")
}
// Nothing more from that burst
select {
case <-changes:
t.Fatal("a burst should report exactly once")
case <-time.After(netChangeSettleWindow):
}
// A change after the window has closed is a separate event and gets its own report.
_, err = unix.Write(kernel, routeMsg(unix.RTM_IFINFO, 8))
require.NoError(t, err)
select {
case <-changes:
case <-time.After(netChangeSettleWindow * 4):
t.Fatal("a later change should report again")
}
// Closing the socket is how the real thing shuts down
require.NoError(t, sock.Close())
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("watchRouteSocket did not return after the socket was closed")
}
}
func TestWatchRouteSocketIgnoresUninterestingMessages(t *testing.T) {
sock, kernel := socketPair(t)
changes := make(chan struct{}, 1)
done := make(chan struct{})
go func() {
watchRouteSocket(test.NewLogger(), sock, changes)
close(done)
}()
for _, msgType := range []uint8{unix.RTM_ADD, unix.RTM_DELETE, unix.RTM_GET, unix.RTM_MISS} {
_, err := unix.Write(kernel, routeMsg(msgType, 8))
require.NoError(t, err)
}
select {
case <-changes:
t.Fatal("route churn alone must not report a change")
case <-time.After(netChangeSettleWindow * 2):
}
require.NoError(t, sock.Close())
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("watchRouteSocket did not return after the socket was closed")
}
}
// TestWatchRouteSocketDropsRatherThanBlocks covers the coalescing send. A reader that is busy rebinding must not
// wedge the watcher, and a second pending "the network moved" tells it nothing new anyway.
func TestWatchRouteSocketDropsRatherThanBlocks(t *testing.T) {
sock, kernel := socketPair(t)
changes := make(chan struct{}, 1)
done := make(chan struct{})
go func() {
watchRouteSocket(test.NewLogger(), sock, changes)
close(done)
}()
// Nobody is reading changes, so after the first report the buffer is full for the rest of this test
for range 3 {
_, err := unix.Write(kernel, routeMsg(unix.RTM_NEWADDR, 8))
require.NoError(t, err)
time.Sleep(netChangeSettleWindow + time.Millisecond*250)
}
// The watcher must still be alive and responsive to a close
require.NoError(t, sock.Close())
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("watchRouteSocket wedged on a full channel")
}
assert.Len(t, changes, 1, "the pending report should have coalesced, not queued")
}
// TestWatchNetworkChangesStopsWithContext covers the detection path against a real routing socket, including that
// cancelling the context closes the channel so a ranging caller falls out of its loop.
func TestWatchNetworkChangesStopsWithContext(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
changes, err := watchNetworkChanges(ctx, test.NewLogger())
require.NoError(t, err)
require.NotNil(t, changes, "darwin should support watching")
drained := make(chan struct{})
go func() {
for range changes {
}
close(drained)
}()
cancel()
select {
case <-drained:
case <-time.After(time.Second * 5):
t.Fatal("cancelling the context should close the changes channel")
}
}
// TestNetworkChangeMonitorStopsWithContext drives the whole monitor against a real routing socket: Start must block
// watching, and cancelling the context (which is all Control does on shutdown, it never stops the monitor directly)
// must return it and clean up the watch goroutines.
func TestNetworkChangeMonitorStopsWithContext(t *testing.T) {
// IgnoreCurrent because other tests in this package leave readers running; we only care about what this test
// leaks itself.
defer goleak.VerifyNone(t, goleak.IgnoreCurrent())
ctx, cancel := context.WithCancel(context.Background())
l := test.NewLogger()
c := config.NewC(l)
require.NoError(t, c.LoadString("listen:\n rebind_on_network_change: true\n"))
m := NewNetworkChangeMonitor(ctx, l, c)
done := make(chan struct{})
go func() {
m.Start(func() {})
close(done)
}()
// Start should be sitting on the routing socket, not have fallen out. If it returned early it either failed to
// watch or no-op'd, both of which we want to catch.
select {
case <-done:
t.Fatal("Start returned instead of watching")
case <-time.After(time.Millisecond * 250):
}
cancel()
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("Start did not return after the context was cancelled")
}
// Starting again after the context is dead must not open anything.
m.Start(func() {})
}
+22
View File
@@ -0,0 +1,22 @@
//go:build !darwin || ios || e2e_testing
// +build !darwin ios e2e_testing
package udp
import (
"context"
"log/slog"
)
// watchNetworkChanges is a no-op outside of darwin.
//
// Darwin is the platform that scopes a udp socket to the interface it came up on, so it is the platform whose socket
// goes stale when the local network changes. Everywhere else Rebind has nothing to do, so there is nothing to watch
// for. iOS is excluded on purpose even though it is darwin: the host app already drives the rebind off NWPathMonitor,
// and two things racing to rebind the same socket is worse than one.
//
// A nil channel means "not supported here", which callers must treat as "do not start a watcher" rather than
// selecting on it, since a receive from a nil channel blocks forever.
func watchNetworkChanges(_ context.Context, _ *slog.Logger) (<-chan struct{}, error) {
return nil, nil
}
+39
View File
@@ -0,0 +1,39 @@
package udp
import (
"context"
"testing"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/test"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func newMonitor(t *testing.T, ctx context.Context, cfg string) *NetworkChangeMonitor {
t.Helper()
l := test.NewLogger()
c := config.NewC(l)
require.NoError(t, c.LoadString(cfg))
return NewNetworkChangeMonitor(ctx, l, c)
}
func TestNetworkChangeMonitorDefaultsOn(t *testing.T) {
// Says nothing about rebinding, so this covers the default.
m := newMonitor(t, context.Background(), "listen:\n host: 0.0.0.0\n")
assert.True(t, m.enabled, "should default to on")
}
func TestNetworkChangeMonitorDisabledIsANoOp(t *testing.T) {
m := newMonitor(t, context.Background(), "listen:\n rebind_on_network_change: false\n")
require.False(t, m.enabled)
// Must return without opening a socket. If it watched anything this would block.
m.Start(func() {})
}
func TestNetworkChangeMonitorNilRebindIsANoOp(t *testing.T) {
// Nothing to rebind, so there is no point watching, on any platform.
m := newMonitor(t, context.Background(), "listen:\n rebind_on_network_change: true\n")
m.Start(nil)
}
-62
View File
@@ -1,62 +0,0 @@
//go:build !android && !e2e_testing
// +build !android,!e2e_testing
package udp
import (
"net"
"syscall"
"unsafe"
"golang.org/x/sys/unix"
)
// rawSendmmsg performs sendmmsg(2) over a syscall.RawConn without
// allocating a closure per call. The struct holds preallocated in/out
// scratch (chunk/sent/errno) and a method-value bound at construction so
// rawConn.Write receives a stable function pointer instead of a fresh
// closure on every send.
type rawSendmmsg struct {
msgs []rawMessage
chunk int
sent int
errno syscall.Errno
callback func(fd uintptr) bool
}
// bind wires r.callback to r.run. Must be called once after r.msgs is set;
// subsequent send calls invoke r.callback without rebinding.
func (r *rawSendmmsg) bind() { r.callback = r.run }
// run is the preallocated callback rawConn.Write invokes. It reads its
// input (r.chunk) and writes its outputs (r.sent, r.errno) through the
// rawSendmmsg fields so the method value does not capture per-call locals
// and therefore does not heap-allocate.
func (r *rawSendmmsg) run(fd uintptr) bool {
r1, _, errno := unix.Syscall6(unix.SYS_SENDMMSG, fd,
uintptr(unsafe.Pointer(&r.msgs[0])), uintptr(r.chunk),
0, 0, 0,
)
if errno == syscall.EAGAIN || errno == syscall.EWOULDBLOCK {
return false
}
r.sent = int(r1)
r.errno = errno
return true
}
// send issues sendmmsg over rc against the first n entries of r.msgs.
// Returns the number of entries the kernel processed and any error;
// matches the original sendmmsg helper's contract.
func (r *rawSendmmsg) send(rc syscall.RawConn, n int) (int, error) {
r.chunk = n
r.sent = 0
r.errno = 0
if err := rc.Write(r.callback); err != nil {
return r.sent, err
}
if r.errno != 0 {
return r.sent, &net.OpError{Op: "sendmmsg", Err: r.errno}
}
return r.sent, nil
}
+16 -10
View File
@@ -140,13 +140,20 @@ func (u *StdConn) WriteTo(b []byte, ap netip.AddrPort) error {
} }
} }
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error { func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
// An un-sendable destination costs its own packet, never the ones behind it in the batch.
// TODO: WriteTo maps EWOULDBLOCK to an error, so a full send buffer
// silently drops the rest of a burst (linux blocks instead). Poll for
// writability on EAGAIN before giving up on the remainder.
written := 0
for i, b := range bufs { for i, b := range bufs {
if err := u.WriteTo(b, addrs[i]); err != nil { if err := u.WriteTo(b, addrs[i]); err == nil {
return err written++
} else {
u.l.Debug("failed to write packet in batch", "udpAddr", addrs[i], "error", err)
} }
} }
return nil return written, nil
} }
func (u *StdConn) LocalAddr() (netip.AddrPort, error) { func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
@@ -188,7 +195,7 @@ func (u *StdConn) ListenOut(r EncReader, flush func()) error {
continue continue
} }
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n], RxMeta{}) r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n:n])
flush() flush()
} }
} }
@@ -197,6 +204,9 @@ func (u *StdConn) SupportsMultipleReaders() bool {
return false return false
} }
// Rebind clears the interface the kernel scoped this socket to, so that sends are routed against the current
// routing table instead of the interface we happened to be on when the socket was created. Darwin pins sockets
// this way on its own, which is what strands us after the underlying network changes.
func (u *StdConn) Rebind() error { func (u *StdConn) Rebind() error {
var err error var err error
if u.isV4 { if u.isV4 {
@@ -205,9 +215,5 @@ func (u *StdConn) Rebind() error {
err = syscall.SetsockoptInt(int(u.sysFd), syscall.IPPROTO_IPV6, syscall.IPV6_BOUND_IF, 0) err = syscall.SetsockoptInt(int(u.sysFd), syscall.IPPROTO_IPV6, syscall.IPV6_BOUND_IF, 0)
} }
if err != nil { return err
u.l.Error("Failed to rebind udp socket", "error", err)
}
return nil
} }
-61
View File
@@ -1,61 +0,0 @@
//go:build linux && !android && !e2e_testing
package udp
import (
"net/netip"
"testing"
)
// TestPlanRunBreaksOnECNChange confirms that two same-destination, same-size
// packets with different outer ECN end up in separate sendmmsg entries (the
// kernel stamps one outer codepoint per entry, so a run that straddled the
// boundary would silently lose information).
func TestPlanRunBreaksOnECNChange(t *testing.T) {
u := &StdConn{gsoSupported: true, maxGSOSegments: 63}
dst := netip.MustParseAddrPort("10.0.0.1:4242")
bufs := [][]byte{
make([]byte, 1200),
make([]byte, 1200),
make([]byte, 1200),
}
addrs := []netip.AddrPort{dst, dst, dst}
t.Run("uniform_ecn_runs_together", func(t *testing.T) {
ecns := []byte{0x02, 0x02, 0x02}
runLen, segSize := u.planRun(bufs, addrs, ecns, 0, 64)
if runLen != 3 {
t.Errorf("runLen=%d want 3 (uniform ECT(0))", runLen)
}
if segSize != 1200 {
t.Errorf("segSize=%d want 1200", segSize)
}
})
t.Run("ecn_change_truncates_run", func(t *testing.T) {
// 0,0,3: first two run together, CE seeds a fresh entry.
ecns := []byte{0x00, 0x00, 0x03}
runLen, _ := u.planRun(bufs, addrs, ecns, 0, 64)
if runLen != 2 {
t.Errorf("runLen=%d want 2 (ECN changes at index 2)", runLen)
}
})
t.Run("nil_ecns_runs_full", func(t *testing.T) {
runLen, _ := u.planRun(bufs, addrs, nil, 0, 64)
if runLen != 3 {
t.Errorf("runLen=%d want 3 (nil ecns means no break)", runLen)
}
})
t.Run("first_ecn_is_singleton", func(t *testing.T) {
// Second packet has different ECN from the first → run halts at 1
// (the first packet alone forms the run).
ecns := []byte{0x00, 0x03, 0x03}
runLen, _ := u.planRun(bufs, addrs, ecns, 0, 64)
if runLen != 1 {
t.Errorf("runLen=%d want 1 (different ECN immediately)", runLen)
}
})
}
+9 -5
View File
@@ -44,13 +44,17 @@ func (u *GenericConn) WriteTo(b []byte, addr netip.AddrPort) error {
return err return err
} }
func (u *GenericConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error { func (u *GenericConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
// An un-sendable destination costs its own packet, never the ones behind it in the batch.
written := 0
for i, b := range bufs { for i, b := range bufs {
if _, err := u.UDPConn.WriteToUDPAddrPort(b, addrs[i]); err != nil { if _, err := u.UDPConn.WriteToUDPAddrPort(b, addrs[i]); err == nil {
return err written++
} else {
u.l.Debug("failed to write packet in batch", "udpAddr", addrs[i], "error", err)
} }
} }
return nil return written, nil
} }
func (u *GenericConn) LocalAddr() (netip.AddrPort, error) { func (u *GenericConn) LocalAddr() (netip.AddrPort, error) {
@@ -102,7 +106,7 @@ func (u *GenericConn) ListenOut(r EncReader, flush func()) error {
continue continue
} }
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n], RxMeta{}) r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n:n])
flush() flush()
} }
} }
+189 -709
View File
File diff suppressed because it is too large Load Diff
-33
View File
@@ -30,39 +30,6 @@ type rawMessage struct {
Len uint32 Len uint32
} }
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
msgs := make([]rawMessage, n)
buffers := make([][]byte, n)
names := make([][]byte, n)
var cmsgs []byte
if cmsgSpace > 0 {
cmsgs = make([]byte, n*cmsgSpace)
}
for i := range msgs {
buffers[i] = make([]byte, bufSize)
names[i] = make([]byte, unix.SizeofSockaddrInet6)
vs := []iovec{
{Base: &buffers[i][0], Len: uint32(len(buffers[i]))},
}
msgs[i].Hdr.Iov = &vs[0]
msgs[i].Hdr.Iovlen = uint32(len(vs))
msgs[i].Hdr.Name = &names[i][0]
msgs[i].Hdr.Namelen = uint32(len(names[i]))
if cmsgSpace > 0 {
msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace]
msgs[i].Hdr.Controllen = uint32(cmsgSpace)
}
}
return msgs, buffers, names, cmsgs
}
func setIovLen(v *iovec, n int) { func setIovLen(v *iovec, n int) {
v.Len = uint32(n) v.Len = uint32(n)
} }
-33
View File
@@ -33,39 +33,6 @@ type rawMessage struct {
Pad0 [4]byte Pad0 [4]byte
} }
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
msgs := make([]rawMessage, n)
buffers := make([][]byte, n)
names := make([][]byte, n)
var cmsgs []byte
if cmsgSpace > 0 {
cmsgs = make([]byte, n*cmsgSpace)
}
for i := range msgs {
buffers[i] = make([]byte, bufSize)
names[i] = make([]byte, unix.SizeofSockaddrInet6)
vs := []iovec{
{Base: &buffers[i][0], Len: uint64(len(buffers[i]))},
}
msgs[i].Hdr.Iov = &vs[0]
msgs[i].Hdr.Iovlen = uint64(len(vs))
msgs[i].Hdr.Name = &names[i][0]
msgs[i].Hdr.Namelen = uint32(len(names[i]))
if cmsgSpace > 0 {
msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace]
msgs[i].Hdr.Controllen = uint64(cmsgSpace)
}
}
return msgs, buffers, names, cmsgs
}
func setIovLen(v *iovec, n int) { func setIovLen(v *iovec, n int) {
v.Len = uint64(n) v.Len = uint64(n)
} }
+558 -103
View File
@@ -3,11 +3,11 @@
package udp package udp
import ( import (
"encoding/binary" "fmt"
"log/slog" "log/slog"
"net" "net"
"net/netip" "net/netip"
"syscall" "slices"
"testing" "testing"
"time" "time"
"unsafe" "unsafe"
@@ -54,39 +54,6 @@ func buildCmsg(level, typ int32, data []byte) []byte {
return buf return buf
} }
// TestParseRecvCmsgOuterECNFamily is the RX half of the dual-stack ECN fix:
// parseRecvCmsg must read the outer ECN from whichever family the kernel
// delivered, not from the socket family. On the default `::` dual-stack bind
// a v4 peer's outer ECN arrives as an IP_TOS cmsg, which the old socket-family
// gate ignored entirely.
func TestParseRecvCmsgOuterECNFamily(t *testing.T) {
tc := make([]byte, 4)
binary.NativeEndian.PutUint32(tc, 0x02)
cases := []struct {
name string
ctrl []byte
want byte
}{
{"ip_tos_ce", buildCmsg(int32(unix.IPPROTO_IP), int32(unix.IP_TOS), []byte{0x03}), 0x03},
{"ip_tos_ect0", buildCmsg(int32(unix.IPPROTO_IP), int32(unix.IP_TOS), []byte{0x02}), 0x02},
{"ipv6_tclass_ect0", buildCmsg(int32(unix.IPPROTO_IPV6), int32(unix.IPV6_TCLASS), tc), 0x02},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
hdr := &msghdr{Control: &c.ctrl[0]}
setMsgControllen(hdr, len(c.ctrl))
gso, ecn := parseRecvCmsg(hdr, false, true)
if gso != 0 {
t.Errorf("gso = %d, want 0 (no UDP_GRO cmsg present)", gso)
}
if ecn != c.want {
t.Errorf("ecn = 0x%02x, want 0x%02x", ecn, c.want)
}
})
}
}
func testLogger() *slog.Logger { func testLogger() *slog.Logger {
return slog.New(slog.DiscardHandler) return slog.New(slog.DiscardHandler)
} }
@@ -123,9 +90,13 @@ func TestWriteBatchBadFamilyDeliversOthers(t *testing.T) {
bufs := [][]byte{[]byte("AAA"), []byte("BBB"), []byte("CCC")} bufs := [][]byte{[]byte("AAA"), []byte("BBB"), []byte("CCC")}
addrs := []netip.AddrPort{good, bad, good} addrs := []netip.AddrPort{good, bad, good}
if err := sender.WriteBatch(bufs, addrs, nil); err != nil { n, err := sender.WriteBatch(bufs, addrs)
if err != nil {
t.Fatalf("WriteBatch returned error, want nil (bad dest should be isolated): %v", err) t.Fatalf("WriteBatch returned error, want nil (bad dest should be isolated): %v", err)
} }
if n != 2 {
t.Errorf("WriteBatch wrote %d packets, want 2 of 3 (the bad-family dest is the only casualty)", n)
}
got := map[string]bool{} got := map[string]bool{}
rx.SetReadDeadline(time.Now().Add(2 * time.Second)) rx.SetReadDeadline(time.Now().Add(2 * time.Second))
@@ -145,12 +116,11 @@ func TestWriteBatchBadFamilyDeliversOthers(t *testing.T) {
} }
} }
// TestWriteBatchOuterTOSToV4Mapped is the TX half of the dual-stack ECN fix, // TestWriteBatchUnreachableDestDeliversOthers is the kernel-rejection twin of
// verified against a live kernel: WriteBatch on the default `::` dual-stack // TestWriteBatchBadFamilyDeliversOthers. A destination the kernel refuses outright (240.0.0.0/4 is reserved, so
// socket, sending to a v4-mapped destination, must stamp the outer ECN via an // the send returns EINVAL) fails its sendmmsg entry; WriteBatch must drop only that entry and still deliver
// IP_TOS cmsg (not IPV6_TCLASS, which the kernel's v4 path ignores) so a v4 // every other packet rather than abandoning the batch at the first failure.
// receiver actually sees it. func TestWriteBatchUnreachableDestDeliversOthers(t *testing.T) {
func TestWriteBatchOuterTOSToV4Mapped(t *testing.T) {
rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
if err != nil { if err != nil {
t.Skipf("cannot open v4 receiver (sandbox?): %v", err) t.Skipf("cannot open v4 receiver (sandbox?): %v", err)
@@ -158,76 +128,561 @@ func TestWriteBatchOuterTOSToV4Mapped(t *testing.T) {
defer rx.Close() defer rx.Close()
rxPort := rx.LocalAddr().(*net.UDPAddr).Port rxPort := rx.LocalAddr().(*net.UDPAddr).Port
// Ask the kernel to deliver the received outer TOS as ancillary data. c, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, false, 1)
rxRaw, err := rx.SyscallConn()
if err != nil { if err != nil {
t.Fatalf("SyscallConn: %v", err) t.Skipf("cannot open v4 sender (sandbox?): %v", err)
}
var soErr error
if err := rxRaw.Control(func(fd uintptr) {
soErr = unix.SetsockoptInt(int(fd), unix.IPPROTO_IP, unix.IP_RECVTOS, 1)
}); err != nil || soErr != nil {
t.Skipf("cannot enable IP_RECVTOS (sandbox/kernel?): ctrl=%v so=%v", err, soErr)
}
c, err := NewListener(testLogger(), netip.IPv6Unspecified(), 0, false, 1)
if err != nil {
t.Skipf("cannot open dual-stack sender (sandbox?): %v", err)
} }
defer c.Close() defer c.Close()
sender := c.(*StdConn) sender := c.(*StdConn)
if sender.isV4 {
t.Skipf("sender came up v4-only; need a dual-stack v6 socket for this test") good := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), uint16(rxPort))
bad := netip.MustParseAddrPort("240.0.0.1:9999") // reserved space, the kernel refuses it
bufs := [][]byte{[]byte("P0"), []byte("P1"), []byte("BAD"), []byte("P3"), []byte("P4")}
addrs := []netip.AddrPort{good, good, bad, good, good}
// The bad destination is reported, but only after every other packet has been attempted.
if _, err := sender.WriteBatch(bufs, addrs); err == nil {
t.Log("WriteBatch returned nil; kernel accepted the reserved address, delivery assertions still apply")
} }
// v4-mapped-in-v6 destination: routed through the kernel's IPv4 path. got := map[string]bool{}
dst := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), uint16(rxPort)) rx.SetReadDeadline(time.Now().Add(2 * time.Second))
const wantECN = byte(0x02) // ECT(0) buf := make([]byte, 64)
for i := 0; i < 4; i++ {
if err := sender.WriteBatch([][]byte{[]byte("tos-probe")}, []netip.AddrPort{dst}, []byte{wantECN}); err != nil { n, _, rerr := rx.ReadFromUDPAddrPort(buf)
t.Fatalf("WriteBatch: %v", err)
}
// Read the datagram plus its ancillary TOS.
rx.SetReadDeadline(time.Now().Add(3 * time.Second))
payload := make([]byte, 128)
oob := make([]byte, 512)
var n, oobn int
var rerr error
if err := rxRaw.Read(func(fd uintptr) bool {
n, oobn, _, _, rerr = unix.Recvmsg(int(fd), payload, oob, 0)
if rerr == syscall.EAGAIN || rerr == syscall.EWOULDBLOCK {
return false
}
return true
}); err != nil {
t.Fatalf("waiting for datagram failed (no delivery?): %v", err)
}
if rerr != nil { if rerr != nil {
t.Fatalf("Recvmsg: %v", rerr) t.Fatalf("expected 4 delivered packets, read #%d failed: %v (got so far: %v)", i+1, rerr, got)
} }
if string(payload[:n]) != "tos-probe" { got[string(buf[:n])] = true
t.Fatalf("payload = %q, want %q", string(payload[:n]), "tos-probe") }
for _, want := range []string{"P0", "P1", "P3", "P4"} {
if !got[want] {
t.Errorf("packet %s was not delivered; delivered set = %v", want, got)
}
}
}
// TestParseRecvCmsgCorruptLenNoPanic: a cmsg Len near max-int used to wrap
// off+clen negative, slip past the bounds check, and drive the walk offset
// negative -- a panic on the next ctrl[off]. The guard must compare Len
// against the remaining bytes instead. Also pins the plain truncated-Len
// cases (too small, larger than the buffer) to a clean early return.
func TestParseRecvCmsgCorruptLenNoPanic(t *testing.T) {
// First cmsg: a valid empty one so the walk advances past off=0
// (off+clen can't overflow while off is still zero).
valid := buildCmsg(int32(unix.SOL_UDP), int32(unix.UDP_GRO), make([]byte, 4))
corrupt := func(lenVal int) []byte {
buf := make([]byte, len(valid)+unix.CmsgSpace(4))
copy(buf, valid)
h := (*unix.Cmsghdr)(unsafe.Pointer(&buf[len(valid)]))
h.Level = int32(unix.IPPROTO_IP)
h.Type = int32(unix.IP_TOS)
setCmsgLen(h, lenVal)
return buf
}
cases := []struct {
name string
ctrl []byte
}{
{"len_near_max_int", corrupt(int(^uint(0)>>1) - 8)},
{"len_too_small", corrupt(unix.SizeofCmsghdr - 1)},
{"len_past_buffer", corrupt(1 << 20)},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
hdr := &msghdr{Control: &c.ctrl[0]}
setMsgControllen(hdr, len(c.ctrl))
gso := parseRecvCmsg(hdr)
// The valid leading UDP_GRO cmsg (payload 0) must still parse;
// the corrupt trailer just ends the walk.
if gso != 0 {
t.Errorf("parseRecvCmsg = %d, want 0", gso)
}
})
}
}
// TestDeliverSegments pins the GRO RX splitting: a kernel-coalesced buffer
// must come back out as the exact pre-coalesce packets -- every boundary
// error here shreds encrypted packets and every decrypt downstream fails.
func TestDeliverSegments(t *testing.T) {
from := netip.MustParseAddrPort("192.0.2.1:4242")
// Spare backing capacity mimics the recvmmsg row a real payload sits in;
// the cap checks below prove none of it leaks to a delivered segment.
pay := func(n int) []byte {
b := make([]byte, n, n+512)
for i := range b {
b[i] = byte(i)
}
return b
}
cases := []struct {
name string
payload []byte
segSize int
wantLens []int
}{
{"no-gro", pay(1400), 0, []int{1400}},
{"negative-segsize", pay(1400), -5, []int{1400}},
{"segsize-equals-payload", pay(1400), 1400, []int{1400}},
{"segsize-past-payload", pay(1400), 2000, []int{1400}},
{"even-split", pay(4200), 1400, []int{1400, 1400, 1400}},
{"short-tail", pay(3000), 1400, []int{1400, 1400, 200}},
{"single-byte-tail", pay(2801), 1400, []int{1400, 1400, 1}},
{"segsize-one", pay(3), 1, []int{1, 1, 1}},
{"empty-payload", pay(0), 1400, []int{0}},
{"max-coalesce", pay(65500), 1372, nil}, // lens derived below
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
wantLens := c.wantLens
if wantLens == nil {
for rem := len(c.payload); rem > 0; rem -= c.segSize {
wantLens = append(wantLens, min(c.segSize, rem))
}
}
var got [][]byte
deliverSegments(func(a netip.AddrPort, seg []byte) {
if a != from {
t.Errorf("from = %v, want %v", a, from)
}
got = append(got, seg)
}, from, c.payload, c.segSize)
if len(got) != len(wantLens) {
t.Fatalf("delivered %d segments, want %d", len(got), len(wantLens))
}
// Segments must tile the payload in order with no gap, overlap,
// or copy: each must alias the payload at the right offset.
off := 0
for i, seg := range got {
if len(seg) != wantLens[i] {
t.Fatalf("segment %d len=%d want %d", i, len(seg), wantLens[i])
}
if cap(seg) != len(seg) {
// EncReader contract: an append into spare capacity would
// scribble into the next segment of the shared row.
t.Errorf("segment %d cap=%d, want %d (capacity must not reach into the row)", i, cap(seg), len(seg))
}
if len(seg) > 0 && &seg[0] != &c.payload[off] {
t.Errorf("segment %d does not alias payload at offset %d", i, off)
}
off += len(seg)
}
if off != len(c.payload) {
t.Errorf("segments cover %d bytes, payload has %d", off, len(c.payload))
}
})
}
}
// newRewindTestWriter builds a batchWriter with no socket: GSO planning on,
// sendFn left for the test to script. fd is invalid on purpose -- any path
// that actually hits the kernel fails loudly.
func newRewindTestWriter() *batchWriter {
w := &batchWriter{fd: -1, isV4: true, l: testLogger()}
w.prepareWriteMessages(MaxWriteBatch)
w.gsoSupported = true
w.maxGSOSegments = 63
return w
}
// capturePrepared decodes n prepared mmsghdr entries beginning at start
// straight from their iovecs -- ground truth, deliberately not the entryEnd
// bookkeeping the resume logic itself relies on. Returns one []byte per
// packed packet, in entry order.
func capturePrepared(w *batchWriter, start, n int) [][]byte {
var out [][]byte
for e := start; e < start+n; e++ {
hdr := &w.msgs[e].Hdr
iovs := unsafe.Slice(hdr.Iov, int(hdr.Iovlen))
for _, iov := range iovs {
b := make([]byte, int(iov.Len))
if iov.Len > 0 {
copy(b, unsafe.Slice(iov.Base, int(iov.Len)))
}
out = append(out, b)
}
}
return out
}
// TestWriteBatchPartialSendRewind drives WriteBatch through scripted
// partial sendmmsg results and asserts the rewind resumes exactly where
// the kernel stopped: every packet on the wire exactly once, in order,
// no duplicate, no loss. This is the hairiest logic in the write path
// and a rewind bug means silent packet duplication or loss under EAGAIN-
// style backpressure.
func TestWriteBatchPartialSendRewind(t *testing.T) {
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
mkBuf := func(tag byte, n int) []byte {
b := make([]byte, n)
for i := range b {
b[i] = tag
}
b[0] = tag // tag identifies the packet uniquely below
return b
}
// Mixed shape: a 3-packet GSO run to A, a lone short packet to A (run
// tail), then two to B. The planner packs this as multiple entries with
// multi-iovec runs, which is what makes the rewind arithmetic hairy.
bufs := [][]byte{
mkBuf(1, 1200), mkBuf(2, 1200), mkBuf(3, 1200), // run to A
mkBuf(4, 600), // short tail to A
mkBuf(5, 900), mkBuf(6, 900), // run to B
}
addrs := []netip.AddrPort{dstA, dstA, dstA, dstA, dstB, dstB}
scripts := [][]int{
{99}, // accept everything first call
{1, 99}, // one entry per call, then the rest
{1, 1, 1, 99}, // strictly one entry per call
{2, 99}, // two entries, then the rest
}
for si, script := range scripts {
t.Run(fmt.Sprintf("script_%d", si), func(t *testing.T) {
w := newRewindTestWriter()
var wire [][]byte
call := 0
w.sendFn = func(start, n int) (int, error) {
accept := n
if call < len(script) && script[call] < n {
accept = script[call]
}
call++
wire = append(wire, capturePrepared(w, start, accept)...)
return accept, nil
}
written, err := w.WriteBatch(bufs, addrs)
if err != nil {
t.Fatalf("WriteBatch: %v", err)
}
if written != len(bufs) {
t.Errorf("written = %d, want %d", written, len(bufs))
}
if len(wire) != len(bufs) {
t.Fatalf("wire got %d packets, want %d (dup or loss in rewind)", len(wire), len(bufs))
}
for i, b := range wire {
if len(b) != len(bufs[i]) || b[0] != bufs[i][0] {
t.Errorf("wire[%d] = tag %d len %d, want tag %d len %d (reorder/dup)",
i, b[0], len(b), bufs[i][0], len(bufs[i]))
}
}
})
}
}
// TestWriteBatchSkipUnroutableRunAccounting: an unroutable destination mid-
// batch is skipped without committing an entry, leaving a hole in the bufs
// index space. The written count must tally packets per sent entry -- the
// index span would count the hole -- across both full and partial sendmmsg
// success.
func TestWriteBatchSkipUnroutableRunAccounting(t *testing.T) {
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
bad := netip.MustParseAddrPort("[2001:db8::1]:9999") // v6 dest, v4 writer
mk := func(tag byte, n int) []byte {
b := make([]byte, n)
b[0] = tag
return b
}
bufs := [][]byte{mk(1, 1200), mk(2, 1200), mk(3, 500), mk(4, 900), mk(5, 900)}
addrs := []netip.AddrPort{dstA, dstA, bad, dstB, dstB}
for si, script := range [][]int{{99}, {1, 99}} {
t.Run(fmt.Sprintf("script_%d", si), func(t *testing.T) {
w := newRewindTestWriter()
var wire [][]byte
call := 0
w.sendFn = func(start, n int) (int, error) {
accept := n
if call < len(script) && script[call] < n {
accept = script[call]
}
call++
wire = append(wire, capturePrepared(w, start, accept)...)
return accept, nil
}
written, err := w.WriteBatch(bufs, addrs)
if err != nil {
t.Fatalf("WriteBatch: %v", err)
}
if written != 4 {
t.Errorf("written = %d, want 4 (the unroutable run is the only casualty)", written)
}
wantTags := []byte{1, 2, 4, 5}
if len(wire) != len(wantTags) {
t.Fatalf("wire got %d packets, want %d (dup or loss around the skip)", len(wire), len(wantTags))
}
for i, b := range wire {
if b[0] != wantTags[i] {
t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i])
}
}
})
}
}
// TestWriteBatchMidChunkRejectResumes: after a partial success, a zero-sent
// error on the FIRST REMAINING entry (done > 0) must drop only that entry's
// run and resume the rest of the chunk in place -- no repacking, no packets
// lost from entries before or after the rejected one.
func TestWriteBatchMidChunkRejectResumes(t *testing.T) {
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
dstC := netip.MustParseAddrPort("127.0.0.3:4242")
mk := func(tag byte, n int) []byte {
b := make([]byte, n)
b[0] = tag
return b
}
// Three entries: a 2-packet GSO run to A, a 2-packet run to B, one to C.
bufs := [][]byte{mk(1, 1200), mk(2, 1200), mk(3, 900), mk(4, 900), mk(5, 600)}
addrs := []netip.AddrPort{dstA, dstA, dstB, dstB, dstC}
w := newRewindTestWriter()
var wire [][]byte
var starts []int
call := 0
w.sendFn = func(start, n int) (int, error) {
starts = append(starts, start)
call++
switch call {
case 1: // accept only entry 0 (the run to A)
wire = append(wire, capturePrepared(w, start, 1)...)
return 1, nil
case 2: // reject entry 1 (the run to B) outright
return -1, &net.OpError{Op: "sendmmsg", Err: unix.EPERM}
default: // accept the rest
wire = append(wire, capturePrepared(w, start, n)...)
return n, nil
}
}
written, err := w.WriteBatch(bufs, addrs)
if err != nil {
t.Fatalf("WriteBatch: %v", err)
}
if written != 3 {
t.Errorf("written = %d, want 3 (B's rejected run is the only casualty)", written)
}
wantTags := []byte{1, 2, 5}
if len(wire) != len(wantTags) {
t.Fatalf("wire got %d packets, want %d (dup or loss around the mid-chunk reject)", len(wire), len(wantTags))
}
for i, b := range wire {
if b[0] != wantTags[i] {
t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i])
}
}
// The resume must reuse the prepared entries: same chunk, advancing
// start offsets, no repack (which would restart at 0 with fresh entries).
if want := []int{0, 1, 2}; !slices.Equal(starts, want) {
t.Errorf("sendFn start offsets = %v, want %v", starts, want)
}
}
// TestWriteBatchMidChunkEIODisablesGSOWithoutDup: an EIO on a GSO entry
// after earlier entries in the chunk already went out must replay ONLY from
// the failed run (replanned as single-packet entries) -- the already-sent
// entries must not be duplicated.
func TestWriteBatchMidChunkEIODisablesGSOWithoutDup(t *testing.T) {
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
mk := func(tag byte, n int) []byte {
b := make([]byte, n)
b[0] = tag
return b
}
// Entry 0: single packet to A. Entry 1: 2-packet GSO run to B.
bufs := [][]byte{mk(1, 600), mk(2, 1200), mk(3, 1200)}
addrs := []netip.AddrPort{dstA, dstB, dstB}
w := newRewindTestWriter()
var wire [][]byte
call := 0
w.sendFn = func(start, n int) (int, error) {
call++
switch call {
case 1: // accept entry 0 only
wire = append(wire, capturePrepared(w, start, 1)...)
return 1, nil
case 2: // EIO on the GSO run to B
return -1, &net.OpError{Op: "sendmmsg", Err: unix.EIO}
default: // replanned single-packet replay
wire = append(wire, capturePrepared(w, start, n)...)
return n, nil
}
}
written, err := w.WriteBatch(bufs, addrs)
if err != nil {
t.Fatalf("WriteBatch: %v", err)
}
if w.gsoSupported {
t.Error("gsoSupported still true after EIO on a GSO entry")
}
if written != len(bufs) {
t.Errorf("written = %d, want %d", written, len(bufs))
}
wantTags := []byte{1, 2, 3}
if len(wire) != len(wantTags) {
t.Fatalf("wire got %d packets, want %d (packet 1 duplicated, or B's run lost)", len(wire), len(wantTags))
}
for i, b := range wire {
if b[0] != wantTags[i] {
t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i])
}
}
}
// TestWriteBatchZeroProgress: sent == 0 with no error must abort with an
// error rather than spin forever replaying the same chunk.
func TestWriteBatchZeroProgress(t *testing.T) {
w := newRewindTestWriter()
w.sendFn = func(start, n int) (int, error) { return 0, nil }
bufs := [][]byte{make([]byte, 100)}
addrs := []netip.AddrPort{netip.MustParseAddrPort("127.0.0.1:4242")}
if _, err := w.WriteBatch(bufs, addrs); err == nil {
t.Fatal("WriteBatch = nil error on zero progress, want error")
}
}
// TestWriteBatchEIODisablesGSOAndReplays pins the runtime GSO give-up: a
// sendmmsg rejected with EIO on a GSO superpacket entry must clear
// gsoSupported and replay the same packets as per-packet entries through
// sendmmsg (keeping batching), not fall back to per-packet sendto.
func TestWriteBatchEIODisablesGSOAndReplays(t *testing.T) {
dst := netip.MustParseAddrPort("127.0.0.1:4242")
bufs := [][]byte{make([]byte, 1200), make([]byte, 1200), make([]byte, 1200)}
addrs := []netip.AddrPort{dst, dst, dst}
w := newRewindTestWriter()
var entryCounts []int
call := 0
w.sendFn = func(start, n int) (int, error) {
entryCounts = append(entryCounts, n)
call++
if call == 1 {
return -1, &net.OpError{Op: "sendmmsg", Err: unix.EIO}
}
return n, nil
}
written, err := w.WriteBatch(bufs, addrs)
if err != nil {
t.Fatalf("WriteBatch: %v", err)
}
if w.gsoSupported {
t.Error("gsoSupported still true after EIO on a GSO entry")
}
if written != len(bufs) {
t.Errorf("written = %d, want %d", written, len(bufs))
}
// First call: one GSO entry carrying the whole run. Replay: one entry
// per packet, still via sendmmsg.
want := []int{1, 3}
if len(entryCounts) != len(want) || entryCounts[0] != want[0] || entryCounts[1] != want[1] {
t.Errorf("sendmmsg entry counts = %v, want %v", entryCounts, want)
}
}
// TestGSOEngagesOnLoopback is the offload smoke test: real sockets, real
// UDP_SEGMENT cmsg, real kernel segmentation over loopback. It asserts
// both that GSO *engaged* (the whole batch left in a single sendmmsg
// entry -- a silent fallback to per-packet entries fails the test) and
// that the kernel carved the superpacket back into the exact original
// datagrams on the receive side. Runs in CI (make test on ubuntu-latest),
// which is what guards against the offload path silently degrading.
func TestGSOEngagesOnLoopback(t *testing.T) {
rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
if err != nil {
t.Fatalf("listen rx: %v", err)
}
defer rx.Close()
dst := rx.LocalAddr().(*net.UDPAddr).AddrPort()
uc, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, false, 8)
if err != nil {
t.Fatalf("NewListener: %v", err)
}
sc := uc.(*StdConn)
defer sc.Close()
if !sc.bw.gsoSupported {
var un unix.Utsname
_ = unix.Uname(&un)
release := string(un.Release[:])
if major, minor := parseRelease(release); major > 4 || (major == 4 && minor >= 18) {
t.Fatalf("kernel %q supports UDP_SEGMENT but the GSO probe failed", release)
}
t.Skipf("kernel %q predates UDP_SEGMENT (4.18)", release)
}
// Spy on the real syscall to count entries per sendmmsg without
// changing what hits the kernel.
var entryCounts []int
real := sc.bw.sendFn
sc.bw.sendFn = func(start, n int) (int, error) {
entryCounts = append(entryCounts, n)
return real(start, n)
}
const numPkts = 8
const pktLen = 1200
bufs := make([][]byte, numPkts)
addrs := make([]netip.AddrPort, numPkts)
for i := range bufs {
bufs[i] = make([]byte, pktLen)
for j := range bufs[i] {
bufs[i][j] = byte(i)
}
addrs[i] = dst
}
written, err := sc.WriteBatch(bufs, addrs)
if err != nil {
t.Fatalf("WriteBatch: %v", err)
}
if written != numPkts {
t.Fatalf("written = %d, want %d", written, numPkts)
}
// GSO engaged means the run went out as ONE sendmmsg entry carrying a
// UDP_SEGMENT superpacket. Per-packet entries mean it silently fell
// back -- exactly the regression this test exists to catch.
if len(entryCounts) != 1 || entryCounts[0] != 1 {
t.Fatalf("sendmmsg entry counts = %v, want [1]: GSO did not engage", entryCounts)
}
// The kernel must deliver the original datagram boundaries and bytes.
_ = rx.SetReadDeadline(time.Now().Add(5 * time.Second))
got := make([]byte, pktLen+1)
for i := 0; i < numPkts; i++ {
n, _, err := rx.ReadFromUDP(got)
if err != nil {
t.Fatalf("rx read %d: %v", i, err)
}
if n != pktLen {
t.Fatalf("rx read %d: len=%d want %d (kernel segmented at wrong boundary)", i, n, pktLen)
}
for j := 0; j < n; j++ {
if got[j] != byte(i) {
t.Fatalf("rx read %d: byte %d = %#x, want %#x", i, j, got[j], byte(i))
} }
cmsgs, err := unix.ParseSocketControlMessage(oob[:oobn])
if err != nil {
t.Fatalf("ParseSocketControlMessage: %v", err)
} }
found := false
var gotTOS byte
for _, m := range cmsgs {
if m.Header.Level == unix.IPPROTO_IP && m.Header.Type == unix.IP_TOS && len(m.Data) >= 1 {
found = true
gotTOS = m.Data[0]
}
}
if !found {
t.Fatalf("no IP_TOS cmsg delivered to v4 receiver — outer ECN did not land (%d cmsgs)", len(cmsgs))
}
if gotTOS&0x03 != wantECN {
t.Errorf("received outer TOS = 0x%02x, want low-2-bits = 0x%02x", gotTOS, wantECN)
} else {
t.Logf("verified: v4 receiver saw outer TOS 0x%02x (ECN=0x%02x) from dual-stack sender", gotTOS, gotTOS&0x03)
} }
} }

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