mirror of
https://github.com/slackhq/nebula.git
synced 2026-10-07 12:57:59 +02:00
the definitive tun offloads branch (#1704)
This commit is contained in:
@@ -0,0 +1,44 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"os"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// blockOn parks the calling goroutine until fd is ready or shutdownFd signals teardown.
|
||||
// (events is POLLIN for reads, POLLOUT for writes)
|
||||
// It builds the pollfd array on the stack every call, so concurrent callers on the same Queue never share Revents storage.
|
||||
//
|
||||
// Returns os.ErrClosed when shutdown was signaled (POLLIN on shutdownFd)
|
||||
// or either fd reported a problem condition (POLLHUP|POLLNVAL|POLLERR).
|
||||
func blockOn(fd, shutdownFd int32, events int16) error {
|
||||
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
||||
pfds := [2]unix.PollFd{
|
||||
{Fd: fd, Events: events},
|
||||
{Fd: shutdownFd, Events: unix.POLLIN},
|
||||
}
|
||||
var err error
|
||||
for {
|
||||
_, err = unix.Poll(pfds[:], -1)
|
||||
if err != unix.EINTR {
|
||||
break
|
||||
}
|
||||
}
|
||||
tunEvents := pfds[0].Revents
|
||||
shutdownEvents := pfds[1].Revents
|
||||
// Check err before trusting the potentially bogus bits we just got.
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
||||
return os.ErrClosed
|
||||
}
|
||||
if tunEvents&problemFlags != 0 {
|
||||
return os.ErrClosed
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"sync/atomic"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
type offloadQueueSet struct {
|
||||
pq []*Offload
|
||||
// pqi is exactly the same as pq, but stored as the interface type
|
||||
pqi []Queue
|
||||
shutdownFd int
|
||||
// usoEnabled is true when newTun successfully negotiated TUN_F_USO4|6 with the kernel.
|
||||
// Queues created by Add inherit this and surface it via Offload.USOSupported so coalescers can gate USO emission.
|
||||
usoEnabled bool
|
||||
closed atomic.Bool
|
||||
// l is handed to each queue for its bad-vnet-header drop logging.
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
// NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do TSO segmentation.
|
||||
// usoEnabled tells downstream queues whether the kernel agreed to deliver/accept GSO_UDP_L4 superpackets.
|
||||
func NewOffloadQueueSet(usoEnabled bool, l *slog.Logger) (QueueSet, error) {
|
||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||
}
|
||||
|
||||
out := &offloadQueueSet{
|
||||
pq: []*Offload{},
|
||||
pqi: []Queue{},
|
||||
shutdownFd: shutdownFd,
|
||||
usoEnabled: usoEnabled,
|
||||
l: l,
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *offloadQueueSet) Queues() []Queue {
|
||||
return c.pqi
|
||||
}
|
||||
|
||||
func (c *offloadQueueSet) Add(fd int) error {
|
||||
if c.closed.Load() {
|
||||
return errors.New("queue set already closed")
|
||||
}
|
||||
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled, c.l)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.pq = append(c.pq, x)
|
||||
c.pqi = append(c.pqi, x)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *offloadQueueSet) wakeForShutdown() error {
|
||||
var buf [8]byte
|
||||
binary.NativeEndian.PutUint64(buf[:], 1)
|
||||
_, err := unix.Write(c.shutdownFd, buf[:])
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *offloadQueueSet) Close() error {
|
||||
if c.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
|
||||
errs := []error{}
|
||||
|
||||
// Signal all readers blocked in poll to wake up and exit.
|
||||
// They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
|
||||
if err := c.wakeForShutdown(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
|
||||
// Close the per-queue tun fds; this also unblocks any in-flight reads.
|
||||
for _, x := range c.pq {
|
||||
if err := x.Close(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Close the shutdown eventfd last: every reader's pollfd set references it,
|
||||
// so it must outlive the wake + per-queue teardown above.
|
||||
if err := unix.Close(c.shutdownFd); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
c.shutdownFd = -1
|
||||
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync/atomic"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
type pollQueueSet struct {
|
||||
pq []*Poll
|
||||
// pqi is exactly the same as pq, but stored as the interface type
|
||||
pqi []Queue
|
||||
shutdownFd int
|
||||
closed atomic.Bool
|
||||
}
|
||||
|
||||
func NewPollQueueSet() (QueueSet, error) {
|
||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||
}
|
||||
|
||||
out := &pollQueueSet{
|
||||
pq: []*Poll{},
|
||||
pqi: []Queue{},
|
||||
shutdownFd: shutdownFd,
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *pollQueueSet) Queues() []Queue {
|
||||
return c.pqi
|
||||
}
|
||||
|
||||
func (c *pollQueueSet) Add(fd int) error {
|
||||
if c.closed.Load() {
|
||||
return errors.New("queue set already closed")
|
||||
}
|
||||
x, err := newPoll(fd, c.shutdownFd)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.pq = append(c.pq, x)
|
||||
c.pqi = append(c.pqi, x)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *pollQueueSet) wakeForShutdown() error {
|
||||
var buf [8]byte
|
||||
binary.NativeEndian.PutUint64(buf[:], 1)
|
||||
_, err := unix.Write(int(c.shutdownFd), buf[:])
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *pollQueueSet) Close() error {
|
||||
if c.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
|
||||
errs := []error{}
|
||||
|
||||
// Signal all readers blocked in poll to wake up and exit.
|
||||
// They observe POLLIN on the shutdown eventfd and return os.ErrClosed.
|
||||
if err := c.wakeForShutdown(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
|
||||
// Close the per-queue tun fds; this also unblocks any in-flight reads.
|
||||
for _, x := range c.pq {
|
||||
if err := x.Close(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Close the shutdown eventfd last: every reader's pollfd set references it,
|
||||
// so it must outlive the wake + per-queue teardown above.
|
||||
if err := unix.Close(c.shutdownFd); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
c.shutdownFd = -1
|
||||
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
|
||||
package tio
|
||||
|
||||
import "testing"
|
||||
|
||||
// fakeBatch stands in for batch.TxBatcher inside the bench — same shape
|
||||
// of pointer-capturing closure that sendInsideMessage builds.
|
||||
type fakeBatch struct{ buf [65536]byte }
|
||||
|
||||
func (b *fakeBatch) Reserve(sz int) []byte { return b.buf[:sz] }
|
||||
func (b *fakeBatch) Commit([]byte) {}
|
||||
|
||||
type fakeHostInfo struct {
|
||||
remoteIndexId uint32
|
||||
counter uint64
|
||||
}
|
||||
type fakeIface struct {
|
||||
rebindCount uint8
|
||||
hi *fakeHostInfo
|
||||
}
|
||||
|
||||
// BenchmarkSegmentSuperpacketAllocsTSO measures allocation per
|
||||
// SegmentSuperpacket call when a closure captures pointer-bearing
|
||||
// receivers — the realistic shape of sendInsideMessage's closure.
|
||||
func BenchmarkSegmentSuperpacketAllocsTSO(b *testing.B) {
|
||||
const mss = 1400
|
||||
const numSeg = 32
|
||||
pkt := buildTSOv6(mss*numSeg, mss)
|
||||
gso := GSOInfo{
|
||||
Size: mss,
|
||||
HdrLen: 60, // 40 (IPv6) + 20 (TCP)
|
||||
CsumStart: 40,
|
||||
Proto: GSOProtoTCP,
|
||||
}
|
||||
p := Packet{Bytes: pkt, GSO: gso}
|
||||
|
||||
hi := &fakeHostInfo{remoteIndexId: 0xdeadbeef}
|
||||
f := &fakeIface{rebindCount: 7, hi: hi}
|
||||
fb := &fakeBatch{}
|
||||
|
||||
// SegmentSuperpacket consumes pkt destructively; refresh from a master
|
||||
// copy each iter (matches the production pattern where every TUN read
|
||||
// hands the segmenter a fresh kernel-supplied buffer).
|
||||
master := append([]byte(nil), pkt...)
|
||||
work := make([]byte, len(pkt))
|
||||
p.Bytes = work
|
||||
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
copy(work, master)
|
||||
err := SegmentSuperpacket(p, func(seg []byte) error {
|
||||
out := fb.Reserve(16 + len(seg) + 16)
|
||||
out[0] = byte(f.rebindCount)
|
||||
out[1] = byte(hi.counter)
|
||||
hi.counter++
|
||||
fb.Commit(out)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
b.Fatalf("SegmentSuperpacket: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
//go:build !linux || android
|
||||
|
||||
package tio
|
||||
|
||||
import "fmt"
|
||||
|
||||
func protoFromGSOType(_ uint8) (GSOProto, error) {
|
||||
return 0, fmt.Errorf("GSO unsupported")
|
||||
}
|
||||
|
||||
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
||||
if pkt.GSO.IsSuperpacket() {
|
||||
return fmt.Errorf("tio: GSO superpacket on platform without segmentation support")
|
||||
}
|
||||
return fn(pkt.Bytes)
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
package tio
|
||||
|
||||
import "io"
|
||||
|
||||
// singleQueue adapts a legacy one-datagram-per-Read source into a Queue.
|
||||
// Read fills a private scratch buffer and returns exactly one Packet whose
|
||||
// Bytes borrow from that buffer, valid only until the next Read, per the Queue contract.
|
||||
// Single-reader like every Queue; Write is exactly as safe for concurrent use as the underlying source's Write.
|
||||
type singleQueue struct {
|
||||
rw io.ReadWriter
|
||||
closer io.Closer // nil: Close is a no-op (the source is shared and owned elsewhere)
|
||||
buf []byte
|
||||
ret [1]Packet
|
||||
}
|
||||
|
||||
// NewSingleQueue wraps a one-datagram-per-Read ReadWriteCloser (a legacy tun device) into a Queue.
|
||||
// bufSize is the per-queue read scratch size and must be at least the largest datagram the source can return.
|
||||
// Close closes rwc.
|
||||
func NewSingleQueue(rwc io.ReadWriteCloser, bufSize int) Queue {
|
||||
return &singleQueue{rw: rwc, closer: rwc, buf: make([]byte, bufSize)}
|
||||
}
|
||||
|
||||
// NewSingleQueueNoClose is NewSingleQueue for a source owned by someone else,
|
||||
// e.g. several queues sharing one device. Close on the returned Queue is a
|
||||
// no-op so one queue can't tear the shared source out from under its
|
||||
// siblings; the owner remains responsible for closing the source itself.
|
||||
func NewSingleQueueNoClose(rw io.ReadWriter, bufSize int) Queue {
|
||||
return &singleQueue{rw: rw, buf: make([]byte, bufSize)}
|
||||
}
|
||||
|
||||
func (q *singleQueue) Read() ([]Packet, error) {
|
||||
n, err := q.rw.Read(q.buf)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
q.ret[0] = Packet{Bytes: q.buf[:n]}
|
||||
return q.ret[:], nil
|
||||
}
|
||||
|
||||
func (q *singleQueue) Write(p []byte) (int, error) {
|
||||
return q.rw.Write(p)
|
||||
}
|
||||
|
||||
func (q *singleQueue) Close() error {
|
||||
if q.closer == nil {
|
||||
return nil
|
||||
}
|
||||
return q.closer.Close()
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
package tio
|
||||
|
||||
import (
|
||||
"io"
|
||||
)
|
||||
|
||||
// QueueSet holds one or many Queue objects and helps close them in an orderly way.
|
||||
type QueueSet interface {
|
||||
io.Closer
|
||||
Queues() []Queue
|
||||
|
||||
// Add takes a tun fd, adds it to the set, and prepares it for use as a Queue.
|
||||
Add(fd int) error
|
||||
}
|
||||
|
||||
// Capabilities advertises which kernel offload features a Queue successfully negotiated.
|
||||
// Callers consult this to decide which coalescers to wire onto the write path.
|
||||
type Capabilities struct {
|
||||
// TSO means the FD was opened with IFF_VNET_HDR and the kernel agreed to TUN_F_TSO4|TSO6,
|
||||
// and WriteGSO with GSOProtoTCP is safe.
|
||||
TSO bool
|
||||
// USO means the kernel additionally agreed to TUN_F_USO4|USO6,
|
||||
// so WriteGSO with GSOProtoUDP is safe. Linux ≥ 6.2.
|
||||
USO bool
|
||||
}
|
||||
|
||||
// Queue is a readable/writable Poll queue.
|
||||
// Concurrency contract: a single read goroutine drives Read; plain Write is safe for concurrent callers;
|
||||
// WriteGSO (on Queues that implement GSOWriter) is single-writer per queue.
|
||||
//
|
||||
// Close on an individual Queue does NOT unblock a Read parked in poll — closing an fd
|
||||
// never wakes its pollers. Orderly teardown goes through the owning QueueSet's Close,
|
||||
// which first signals a shared shutdown eventfd every reader polls alongside its own fd.
|
||||
// That eventfd is a set-wide kill switch: once signaled, every Queue in the set returns
|
||||
// os.ErrClosed from Read, so it cannot be used to stop a single Queue.
|
||||
type Queue interface {
|
||||
io.Closer
|
||||
|
||||
// Read returns one or more packets.
|
||||
// The returned Packet.Bytes slices are borrowed from the Queue's internal buffer and are only valid
|
||||
// until the next Read or Close on this Queue.
|
||||
// A Packet may carry a GSO/USO superpacket (see GSOInfo)
|
||||
// Single-reader only: not safe for concurrent Reads (it reuses per-queue rx scratch each call).
|
||||
Read() ([]Packet, error)
|
||||
|
||||
// Write emits a single packet on the plaintext (outside→inside) delivery path.
|
||||
// Safe for concurrent use.
|
||||
Write(p []byte) (int, error)
|
||||
}
|
||||
|
||||
// Packet is the unit Queue.Read returns.
|
||||
// Bytes points into the queue's internal buffer and is only valid until the next Read or Close on the queue that produced it.
|
||||
// GSO is the zero value for an already-segmented IP datagram;
|
||||
// when non-zero it describes a kernel-supplied TSO/USO superpacket the caller must segment before consuming.
|
||||
type Packet struct {
|
||||
Bytes []byte
|
||||
GSO GSOInfo
|
||||
}
|
||||
|
||||
// GSOInfo describes a kernel-supplied superpacket sitting in Packet.Bytes.
|
||||
// The zero value means Bytes is one regular IP datagram and no segmentation is required.
|
||||
type GSOInfo struct {
|
||||
// Size is the GSO segment size: max payload bytes per segment
|
||||
// (== TCP MSS for TSO, == UDP payload chunk for USO). Zero means not a superpacket.
|
||||
Size uint16
|
||||
// HdrLen is the total L3+L4 header length within Bytes (already corrected via correctHdrLen, so safe to slice on).
|
||||
HdrLen uint16
|
||||
// CsumStart is the L4 header offset inside Bytes (== L3 header length).
|
||||
CsumStart uint16
|
||||
// Proto picks the L4 protocol (TCP or UDP) so the segmenter knows which checksum/header layout to apply.
|
||||
Proto GSOProto
|
||||
}
|
||||
|
||||
// IsSuperpacket reports whether g describes a multi-segment GSO/USO
|
||||
// superpacket that needs segmentation before its bytes can be encrypted and sent on the wire.
|
||||
func (g GSOInfo) IsSuperpacket() bool { return g.Size > 0 }
|
||||
|
||||
// Clone returns a Packet whose Bytes is a freshly allocated copy of p.Bytes,
|
||||
// safe to retain past the next Read or Close on the originating Queue.
|
||||
// GSO metadata is copied verbatim.
|
||||
// Use this only when a caller needs the data to outlive the borrowed-slice contract.
|
||||
func (p Packet) Clone() Packet {
|
||||
if p.Bytes == nil {
|
||||
return p
|
||||
}
|
||||
cp := make([]byte, len(p.Bytes))
|
||||
copy(cp, p.Bytes)
|
||||
return Packet{Bytes: cp, GSO: p.GSO}
|
||||
}
|
||||
|
||||
// CapsProvider is an optional interface implemented by Queues that negotiate kernel offload features at open time.
|
||||
// Callers pick a write-path coalescer based on the result.
|
||||
// Queues that don't implement it are treated as having no offload capability.
|
||||
type CapsProvider interface {
|
||||
Capabilities() Capabilities
|
||||
}
|
||||
|
||||
// GSOProto selects the L4 protocol for a GSO superpacket.
|
||||
// Determines which VIRTIO_NET_HDR_GSO_* type the writer stamps and which checksum offset
|
||||
// inside the transport header virtio NEEDS_CSUM expects.
|
||||
type GSOProto uint8
|
||||
|
||||
const (
|
||||
GSOProtoUnknown GSOProto = iota
|
||||
GSOProtoTCP
|
||||
GSOProtoUDP
|
||||
)
|
||||
|
||||
// GSOWriter is implemented by Queues that can emit a TCP or UDP superpacket
|
||||
// assembled from a header prefix plus one or more borrowed payload fragments,
|
||||
// in a single vectored write (writev with a leading virtio_net_hdr).
|
||||
// This lets the coalescer avoid copying payload bytes between the caller's decrypt buffer and the TUN.
|
||||
// Backends without GSO support do not implement this interface and coalescing is skipped.
|
||||
//
|
||||
// hdr contains the IPv4/IPv6 header prefix (mutable: callers will have filled in total length and IP csum).
|
||||
// transportHdr is the TCP or UDP header
|
||||
// (mutable: the L4 checksum field must hold the pseudo-header partial, single-fold not inverted, per virtio NEEDS_CSUM semantics).
|
||||
// pays are non-overlapping payload fragments whose concatenation is the full superpacket payload.
|
||||
// They are read-only from the writer's perspective and must remain valid until the call returns.
|
||||
// Every segment in pays except possibly the last must be exactly the same size.
|
||||
// proto picks the L4 protocol so the writer knows which gsoType / CsumOffset to set.
|
||||
//
|
||||
// Callers should also consult CapsProvider (via SupportsGSO) for the per-protocol negotiated capability:
|
||||
// USO may not have been negotiated even when TSO was.
|
||||
type GSOWriter interface {
|
||||
io.Writer
|
||||
CapsProvider
|
||||
WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error
|
||||
}
|
||||
|
||||
// SupportsGSO reports whether w implements GSOWriter and the underlying
|
||||
// queue advertises the negotiated capability for `want`.
|
||||
func SupportsGSO(w io.Writer, want GSOProto) (GSOWriter, bool) {
|
||||
gw, ok := w.(GSOWriter)
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
caps := gw.Capabilities()
|
||||
switch want {
|
||||
case GSOProtoTCP:
|
||||
return gw, caps.TSO
|
||||
case GSOProtoUDP:
|
||||
return gw, caps.USO
|
||||
default:
|
||||
return gw, false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,417 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"os"
|
||||
"sync/atomic"
|
||||
"syscall"
|
||||
"unsafe"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio/virtio"
|
||||
)
|
||||
|
||||
const maxSuperpacketLen = 65535
|
||||
|
||||
// tunRxBufSize is the per-Read worst-case footprint inside rxBuf: one kernel-supplied packet body, which is at most ~64 KiB.
|
||||
// Segmentation happens at encrypt time on a per-routine MTU-sized scratch
|
||||
// (see SegmentSuperpacket), so rxBuf only holds raw kernel-supplied bytes.
|
||||
// We round up to give margin for the drain headroom check below.
|
||||
const tunRxBufSize = 64 * 1024
|
||||
|
||||
// tunRxBufCap is the total size we allocate for the per-reader rx buffer.
|
||||
// Each drain iteration consumes up to tunRxBufSize of headroom for the kernel-supplied bytes.
|
||||
// Sized to eight such iterations so a single poll wake can drain several TSO/USO superpackets under bulk load,
|
||||
// amortizing the wake and giving the sendmmsg planner longer same-destination runs.
|
||||
// Hold latency stays bounded because listenIn flushes its send batch incrementally rather than only at end-of-drain.
|
||||
const tunRxBufCap = tunRxBufSize * 8
|
||||
|
||||
// tunDrainCap caps how many packets a single Read will accumulate via the post-wake drain loop.
|
||||
// Sized to soak up a burst of small ACKs while bounding how much work a single caller holds before handing off.
|
||||
const tunDrainCap = 64
|
||||
|
||||
// gsoMaxIovs caps the iovec budget WriteGSO assembles per call:
|
||||
// 3 fixed entries (virtio_net_hdr, IP hdr, transport hdr), plus up to gsoMaxIovs-3 payload fragments.
|
||||
// Sized comfortably above the typical kernel GSO segment cap (Linux UDP_GRO is 64)
|
||||
// so realistic coalesced bursts never touch the limit.
|
||||
// iovecs are tiny (16 bytes), so the entire scratch is 4 KiB.
|
||||
// WriteGSO returns an error rather than reallocating when a caller exceeds this budget.
|
||||
const gsoMaxIovs = 256
|
||||
|
||||
// validVnetHdr is the 10-byte virtio_net_hdr we prepend to every non-GSO TUN write.
|
||||
// Only flag set is VIRTIO_NET_HDR_F_DATA_VALID. Note the tun write path
|
||||
// (__virtio_net_hdr_to_skb) ignores this bit — only the virtio-net driver's RX
|
||||
// helper honors it — so packets land CHECKSUM_NONE and the stack verifies the
|
||||
// L4 checksum anyway. What matters here is what the header does NOT say:
|
||||
// no NEEDS_CSUM, so the kernel is never asked to finish a checksum.
|
||||
// All packets that reach the plain Write paths already carry a valid L4 checksum.
|
||||
var validVnetHdr = [virtio.Size]byte{unix.VIRTIO_NET_HDR_F_DATA_VALID}
|
||||
|
||||
// Offload wraps a TUN file descriptor with poll-based reads. The FD provided will be changed to non-blocking.
|
||||
// A shared eventfd allows Close to wake all readers blocked in poll.
|
||||
//
|
||||
// Field order is deliberate: the read-mostly fds and the writer-owned GSO scratch fill
|
||||
// the first cache line, and the state the reader mutates per packet (rxOff, pending,
|
||||
// readIovs) all sits after it, so per-packet reader stores never invalidate the line
|
||||
// concurrent Write callers load fd from.
|
||||
type Offload struct {
|
||||
fd int
|
||||
shutdownFd int
|
||||
// usoEnabled records whether the kernel agreed to TUN_F_USO* on this FD,
|
||||
// so writers can decide whether emitting GSO_UDP_L4 superpackets is safe.
|
||||
usoEnabled bool
|
||||
closed atomic.Bool
|
||||
|
||||
// gsoHdrBuf is a per-queue 10-byte scratch for the virtio_net_hdr emitted
|
||||
// by WriteGSO. Kept separate from the read-only package-level validVnetHdr
|
||||
// so non-GSO Writes can ship that constant directly while WriteGSO
|
||||
// rewrites this scratch on every call.
|
||||
gsoHdrBuf [virtio.Size]byte
|
||||
// gsoIovs is the writev iovec scratch for WriteGSO. Pre-sized to
|
||||
// gsoMaxIovs at construction; never grown. WriteGSO returns an error
|
||||
// (and drops the call) if a caller hands it more fragments than fit.
|
||||
gsoIovs []unix.Iovec
|
||||
|
||||
rxBuf []byte // backing store for kernel-handed packets read this drain
|
||||
rxOff int // cursor into rxBuf for the current Read drain
|
||||
pending []Packet // packets returned from the most recent Read
|
||||
|
||||
// readVnetScratch holds the 10-byte virtio_net_hdr split off the front of
|
||||
// every TUN read via readv(2). Decoupling the header from the packet body
|
||||
// lets us read the body directly into rxBuf at the current rxOff with
|
||||
// no userspace copy on the GSO_NONE fast path.
|
||||
readVnetScratch [virtio.Size]byte
|
||||
// readIovs is the readv(2) iovec scratch wired once at construction,
|
||||
// iovec[0] points at readVnetScratch
|
||||
// iovec[1].Base/Len is updated per read to address the current rxBuf slot.
|
||||
readIovs [2]unix.Iovec
|
||||
|
||||
// l is only consulted on the rare bad-vnet-header drop path; it lives
|
||||
// after the hot state on purpose. May be nil (tests); drops go unlogged then.
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
func newOffload(fd int, shutdownFd int, usoEnabled bool, l *slog.Logger) (*Offload, error) {
|
||||
if err := unix.SetNonblock(fd, true); err != nil {
|
||||
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
||||
}
|
||||
|
||||
out := &Offload{
|
||||
fd: fd,
|
||||
shutdownFd: shutdownFd,
|
||||
usoEnabled: usoEnabled,
|
||||
closed: atomic.Bool{},
|
||||
l: l,
|
||||
|
||||
rxBuf: make([]byte, tunRxBufCap),
|
||||
gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs),
|
||||
}
|
||||
|
||||
out.gsoIovs[0].Base = &out.gsoHdrBuf[0]
|
||||
out.gsoIovs[0].SetLen(virtio.Size)
|
||||
|
||||
// readIovs[0] is wired once to the virtio_net_hdr scratch; per-read we
|
||||
// only repoint readIovs[1] at the next rxBuf slot (see readPacket).
|
||||
out.readIovs[0].Base = &out.readVnetScratch[0]
|
||||
out.readIovs[0].SetLen(virtio.Size)
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (r *Offload) blockOnRead() error {
|
||||
return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLIN)
|
||||
}
|
||||
|
||||
func (r *Offload) blockOnWrite() error {
|
||||
return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLOUT)
|
||||
}
|
||||
|
||||
// readPacket issues a single readv(2), splitting the virtio_net_hdr off into readVnetScratch
|
||||
// and reading the packet body directly into rxBuf at the current rxOff.
|
||||
// Returns the body length (zero virtio header bytes, just the IP packet/superpacket).
|
||||
// block controls whether EAGAIN is retried via poll: the initial read of a drain blocks; subsequent drain reads do not.
|
||||
func (r *Offload) readPacket(block bool) (int, error) {
|
||||
for {
|
||||
r.readIovs[1].Base = &r.rxBuf[r.rxOff]
|
||||
r.readIovs[1].SetLen(len(r.rxBuf) - r.rxOff)
|
||||
n, _, errno := syscall.Syscall(unix.SYS_READV, uintptr(r.fd), uintptr(unsafe.Pointer(&r.readIovs[0])), uintptr(len(r.readIovs)))
|
||||
if errno == 0 {
|
||||
if int(n) < virtio.Size {
|
||||
return 0, fmt.Errorf("tun read shorter than virtio_net_hdr: %d bytes", n)
|
||||
}
|
||||
return int(n) - virtio.Size, nil
|
||||
}
|
||||
if errno == unix.EAGAIN {
|
||||
if !block {
|
||||
return 0, errno
|
||||
}
|
||||
if err := r.blockOnRead(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
continue
|
||||
}
|
||||
if errno == unix.EINTR {
|
||||
continue
|
||||
}
|
||||
if errno == unix.EBADF {
|
||||
return 0, os.ErrClosed
|
||||
}
|
||||
return 0, errno
|
||||
}
|
||||
}
|
||||
|
||||
// Read returns one or more packets from the tun.
|
||||
// Each Packet either carries a single ready-to-use IP datagram (GSO zero) or a TSO/USO superpacket plus the GSOInfo a caller needs to segment it (see SegmentSuperpacket).
|
||||
// The first read blocks via poll; once the fd is known readable we drain additional packets non-blocking until:
|
||||
// - the kernel queue is empty (EAGAIN)
|
||||
// - we've collected tunDrainCap packets,
|
||||
// - or we're out of rxBuf headroom.
|
||||
//
|
||||
// This amortizes the poll wake over bursts of small packets (e.g. TCP ACKs).
|
||||
// Packet.Bytes slices point into the Offload's internal buffer and are only valid until the next Read or Close on this Queue.
|
||||
func (r *Offload) Read() ([]Packet, error) {
|
||||
r.pending = r.pending[:0]
|
||||
r.rxOff = 0
|
||||
|
||||
// Initial (blocking) read.
|
||||
// Retry on decode errors so a single bad packet does not stall the reader.
|
||||
for {
|
||||
n, err := r.readPacket(true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := r.decodeRead(n); err != nil {
|
||||
// Drop and read again. A bad packet should not kill the reader,
|
||||
// but a systematic decode failure must not be invisible either.
|
||||
r.logDroppedRead(err)
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
|
||||
// Drain: non-blocking reads until the kernel queue is empty, the drain
|
||||
// cap is reached, or rxBuf no longer has room for another worst-case
|
||||
// kernel-supplied packet (tunRxBufSize).
|
||||
for len(r.pending) < tunDrainCap && tunRxBufCap-r.rxOff >= tunRxBufSize {
|
||||
n, err := r.readPacket(false)
|
||||
if err != nil {
|
||||
// EAGAIN / EINTR / anything else: stop draining. We already
|
||||
// have a valid batch from the first read.
|
||||
break
|
||||
}
|
||||
if n <= 0 {
|
||||
break
|
||||
}
|
||||
if err := r.decodeRead(n); err != nil {
|
||||
// Drop this packet and stop the drain; we'd rather hand off
|
||||
// what we have than keep spinning here.
|
||||
r.logDroppedRead(err)
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
return r.pending, nil
|
||||
}
|
||||
|
||||
// logDroppedRead reports a tun packet dropped for a bad/unsupported virtio
|
||||
// header. Debug-gated so the happy path never pays for attribute assembly.
|
||||
func (r *Offload) logDroppedRead(err error) {
|
||||
if r.l != nil && r.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
r.l.Debug("dropping tun packet with bad virtio header", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// decodeRead processes the packet sitting in rxBuf at rxOff (length pktLen).
|
||||
// The bytes stay in rxBuf:
|
||||
// - for GSO_NONE we slice them as a regular IP datagram (running finishChecksum if NEEDS_CSUM is set);
|
||||
// - for TSO/USO superpackets we attach the corrected GSO metadata, so the caller can segment lazily at encrypt time.
|
||||
//
|
||||
// rxOff advances by pktLen on success
|
||||
func (r *Offload) decodeRead(pktLen int) error {
|
||||
if pktLen <= 0 {
|
||||
return fmt.Errorf("short tun read: %d", pktLen)
|
||||
}
|
||||
var hdr virtio.Hdr
|
||||
hdr.Decode(r.readVnetScratch[:])
|
||||
|
||||
body := r.rxBuf[r.rxOff : r.rxOff+pktLen]
|
||||
|
||||
if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
|
||||
if err := virtio.FinishChecksum(body, hdr); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
r.pending = append(r.pending, Packet{Bytes: body})
|
||||
r.rxOff += pktLen
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := virtio.CheckValid(body, hdr); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := virtio.CorrectHdrLen(body, &hdr); err != nil {
|
||||
return err
|
||||
}
|
||||
proto, err := protoFromGSOType(hdr.GSOType())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
r.pending = append(r.pending, Packet{
|
||||
Bytes: body,
|
||||
GSO: GSOInfo{
|
||||
Size: hdr.GSOSize,
|
||||
HdrLen: hdr.HdrLen,
|
||||
CsumStart: hdr.CsumStart,
|
||||
Proto: proto,
|
||||
},
|
||||
})
|
||||
r.rxOff += pktLen
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Offload) Write(buf []byte) (int, error) {
|
||||
if len(buf) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
iovs := [2]unix.Iovec{
|
||||
{Base: &validVnetHdr[0]},
|
||||
{Base: &buf[0]},
|
||||
}
|
||||
iovs[0].SetLen(virtio.Size)
|
||||
iovs[1].SetLen(len(buf))
|
||||
return r.rawWrite(unsafe.Slice(&iovs[0], 2))
|
||||
}
|
||||
|
||||
func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
|
||||
for {
|
||||
n, _, errno := syscall.Syscall(unix.SYS_WRITEV, uintptr(r.fd), uintptr(unsafe.Pointer(&iovs[0])), uintptr(len(iovs)))
|
||||
if errno == 0 {
|
||||
if int(n) < virtio.Size {
|
||||
return 0, io.ErrShortWrite
|
||||
}
|
||||
return int(n) - virtio.Size, nil
|
||||
}
|
||||
if errno == unix.EAGAIN {
|
||||
if err := r.blockOnWrite(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
continue
|
||||
}
|
||||
if errno == unix.EINTR {
|
||||
continue
|
||||
}
|
||||
if errno == unix.EBADF {
|
||||
return 0, os.ErrClosed
|
||||
}
|
||||
return 0, errno
|
||||
}
|
||||
}
|
||||
|
||||
// Capabilities reports the offload features negotiated for this Queue. TSO
|
||||
// is always true for Offload (we only construct it on IFF_VNET_HDR FDs);
|
||||
// USO is true only when the kernel agreed to TUN_F_USO4|6 at open time (Linux ≥ 6.2).
|
||||
func (r *Offload) Capabilities() Capabilities {
|
||||
return Capabilities{TSO: true, USO: r.usoEnabled}
|
||||
}
|
||||
|
||||
func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error {
|
||||
if len(pays) == 0 {
|
||||
// There are no payload fragments. There is nothing to send.
|
||||
return nil
|
||||
}
|
||||
var csumOff uint16 // csumOff is the offset of the L4 checksum field in transportHdr
|
||||
switch proto {
|
||||
case GSOProtoUDP:
|
||||
csumOff = 6
|
||||
case GSOProtoTCP:
|
||||
csumOff = 16
|
||||
default:
|
||||
return fmt.Errorf("unknown GSO proto: %d", proto)
|
||||
}
|
||||
// Incorrect geometry must cause an error, not a silent drop.
|
||||
// No sane packet should ever make it inside this branch.
|
||||
if len(hdr) == 0 || len(transportHdr) < int(csumOff)+2 {
|
||||
return fmt.Errorf("tio: WriteGSO header too short: ip=%d transport=%d (csum field at %d)", len(hdr), len(transportHdr), csumOff)
|
||||
}
|
||||
// Make the iovec array: [virtio_hdr, hdr, transportHdr, pays...].
|
||||
// The constructor attaches r.gsoIovs[0] to gsoHdrBuf. That entry does not change.
|
||||
need := 3 + len(pays)
|
||||
if need > cap(r.gsoIovs) {
|
||||
return fmt.Errorf("tio: WriteGSO needs %d iovecs but cap is %d", need, cap(r.gsoIovs))
|
||||
}
|
||||
r.gsoIovs = r.gsoIovs[:need]
|
||||
r.gsoIovs[1].Base = &hdr[0]
|
||||
r.gsoIovs[1].SetLen(len(hdr))
|
||||
r.gsoIovs[2].Base = &transportHdr[0]
|
||||
r.gsoIovs[2].SetLen(len(transportHdr))
|
||||
|
||||
segSize := len(pays[0])
|
||||
total := len(hdr) + len(transportHdr)
|
||||
for i, p := range pays {
|
||||
if len(p) == 0 {
|
||||
// The coalescers route zero-payload packets down the non-GSO path,
|
||||
// so an empty fragment means the caller's accounting is broken.
|
||||
return fmt.Errorf("tio: WriteGSO empty payload fragment %d of %d", i, len(pays))
|
||||
} else if len(p) > segSize || (len(p) < segSize && i != len(pays)-1) {
|
||||
// all segments must be the same size, except for the last one
|
||||
return fmt.Errorf("tio: WriteGSO fragment %d is %dB, want %dB segments (only the last may be shorter)", i, len(p), segSize)
|
||||
}
|
||||
total += len(p)
|
||||
r.gsoIovs[3+i].Base = &p[0]
|
||||
r.gsoIovs[3+i].SetLen(len(p))
|
||||
}
|
||||
// This check keeps `total` in the uint16 range. Anything larger would wrap around and cause trouble.
|
||||
if total > maxSuperpacketLen {
|
||||
return fmt.Errorf("tio: WriteGSO superpacket %dB exceeds %d", total, maxSuperpacketLen)
|
||||
}
|
||||
|
||||
// A single segment ships as a plain checksummed packet (GSO_NONE, size 0).
|
||||
// Multiple segments carry the real GSO type and segSize, which the loop
|
||||
// above verified is the size of every fragment except possibly the last.
|
||||
gsoType := uint8(unix.VIRTIO_NET_HDR_GSO_NONE)
|
||||
if len(pays) > 1 {
|
||||
gsoType = gsoTypeFromProto(proto, hdr[0]>>4)
|
||||
if gsoType == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||
// gsoTypeFromProto only yields GSO_NONE for a bogus IP version nibble.
|
||||
// A multi-segment superpacket must carry a real GSO type, or the kernel would deliver it as a single jumbo packet.
|
||||
return fmt.Errorf("tio: WriteGSO IP version %d is not GSO-capable", hdr[0]>>4)
|
||||
}
|
||||
}
|
||||
var gsoSize uint16
|
||||
if gsoType != unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||
gsoSize = uint16(segSize)
|
||||
}
|
||||
virtio.EncodeHeader(
|
||||
r.gsoHdrBuf[:],
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
gsoType, /*gsoType*/
|
||||
uint16(len(hdr)+len(transportHdr)), /*hdrLen*/
|
||||
gsoSize, /*gsoSize*/
|
||||
uint16(len(hdr)), /*csumStart*/
|
||||
csumOff, /*csumOffset*/
|
||||
)
|
||||
|
||||
_, err := r.rawWrite(r.gsoIovs)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *Offload) Close() error {
|
||||
if r.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
|
||||
// shutdownFd is owned by the container, so we should not close it
|
||||
// Close the underlying fd but do NOT null r.fd: a reader may still be loading it in readPacket, and mutating the field would race that load.
|
||||
// That reader gets EBADF -> os.ErrClosed on its next syscall. A reader already parked in
|
||||
// poll is NOT woken by this close; only the QueueSet's shutdown eventfd wake does that (see Queue.Close docs).
|
||||
// closed.Swap already guarantees we only close once.
|
||||
return unix.Close(r.fd)
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"sync/atomic"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
type Poll struct {
|
||||
fd int
|
||||
shutdownFd int
|
||||
closed atomic.Bool
|
||||
|
||||
readBuf []byte
|
||||
batchRet [1]Packet
|
||||
}
|
||||
|
||||
// newPoll wraps an existing tun fd.
|
||||
// On failure it does NOT close fd: the caller owns fd and is the sole closer
|
||||
// (see pollQueueSet.Add callers in overlay/tun_linux.go, which unix.Close on Add error).
|
||||
// This matches the newOffload convention and keeps closes at exactly one on every path.
|
||||
func newPoll(fd int, shutdownFd int) (*Poll, error) {
|
||||
if err := unix.SetNonblock(fd, true); err != nil {
|
||||
return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err)
|
||||
}
|
||||
|
||||
out := &Poll{
|
||||
fd: fd,
|
||||
shutdownFd: shutdownFd,
|
||||
readBuf: make([]byte, 65535), // largest possible size Linux permits
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// blockOnRead waits until the Poll fd is readable or shutdown has been signaled.
|
||||
// Returns os.ErrClosed if Close was called.
|
||||
func (t *Poll) blockOnRead() error {
|
||||
return blockOn(int32(t.fd), int32(t.shutdownFd), unix.POLLIN)
|
||||
}
|
||||
|
||||
func (t *Poll) blockOnWrite() error {
|
||||
return blockOn(int32(t.fd), int32(t.shutdownFd), unix.POLLOUT)
|
||||
}
|
||||
|
||||
// TODO: port Offload's post-wake drain loop here so one poll wake amortizes
|
||||
// over a burst (up to tunDrainCap packets) instead of paying a syscall and a
|
||||
// wake per packet. Hosts on the TUNSETOFFLOAD-failure fallback or a tun.fd
|
||||
// config currently lose that batching. blockOn and the EAGAIN plumbing are
|
||||
// already shared; kept one-packet-per-Read for now to preserve behavior.
|
||||
func (t *Poll) Read() ([]Packet, error) {
|
||||
n, err := t.readOne(t.readBuf)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
t.batchRet[0] = Packet{Bytes: t.readBuf[:n]}
|
||||
return t.batchRet[:], nil
|
||||
}
|
||||
|
||||
func (t *Poll) readOne(to []byte) (int, error) {
|
||||
for {
|
||||
n, errno := unix.Read(t.fd, to)
|
||||
if errno == nil {
|
||||
return n, nil
|
||||
}
|
||||
switch errno {
|
||||
case unix.EAGAIN:
|
||||
if err := t.blockOnRead(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
case unix.EINTR:
|
||||
// retry
|
||||
case unix.EBADF:
|
||||
return 0, os.ErrClosed
|
||||
default:
|
||||
return 0, errno
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Write is safe for concurrent use
|
||||
func (t *Poll) Write(from []byte) (int, error) {
|
||||
for {
|
||||
n, errno := unix.Write(t.fd, from)
|
||||
if errno == nil {
|
||||
return n, nil
|
||||
}
|
||||
switch errno {
|
||||
case unix.EAGAIN:
|
||||
if err := t.blockOnWrite(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
case unix.EINTR:
|
||||
// retry
|
||||
case unix.EBADF:
|
||||
return 0, os.ErrClosed
|
||||
default:
|
||||
return 0, errno
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (t *Poll) Close() error {
|
||||
if t.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
|
||||
// shutdownFd is owned by the container, so we should not close it
|
||||
// Close the underlying fd but do NOT null t.fd: a reader may still be loading it in readOne, and mutating the field would race that load.
|
||||
// That reader gets EBADF -> os.ErrClosed on its next syscall. A reader already parked in
|
||||
// poll is NOT woken by this close; only the QueueSet's shutdown eventfd wake does that (see Queue.Close docs).
|
||||
// closed.Swap already guarantees we only close once.
|
||||
return unix.Close(t.fd)
|
||||
}
|
||||
@@ -0,0 +1,228 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
// +build linux,!android,!e2e_testing
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"log/slog"
|
||||
"os"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// newReadPipe returns a read fd. The matching write fd is registered for cleanup.
|
||||
// The caller takes ownership of the read fd (pass it into a QueueSet).
|
||||
func newReadPipe(t *testing.T) int {
|
||||
t.Helper()
|
||||
var fds [2]int
|
||||
if err := unix.Pipe2(fds[:], unix.O_CLOEXEC); err != nil {
|
||||
t.Fatalf("pipe2: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = unix.Close(fds[1]) })
|
||||
return fds[0]
|
||||
}
|
||||
|
||||
func TestPoll_WakeForShutdown_WakesFriends(t *testing.T) {
|
||||
pipe1 := newReadPipe(t)
|
||||
pipe2 := newReadPipe(t)
|
||||
parent, err := NewPollQueueSet()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, parent.Add(pipe1))
|
||||
require.NoError(t, parent.Add(pipe2))
|
||||
t.Cleanup(func() {
|
||||
_ = unix.Close(pipe1)
|
||||
_ = unix.Close(pipe2)
|
||||
})
|
||||
|
||||
readers := parent.Queues()
|
||||
errs := make([]error, len(readers))
|
||||
var wg sync.WaitGroup
|
||||
for i, r := range readers {
|
||||
wg.Add(1)
|
||||
go func(i int, r Queue) {
|
||||
defer wg.Done()
|
||||
_, errs[i] = r.Read()
|
||||
}(i, r)
|
||||
}
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
if err := parent.Close(); err != nil {
|
||||
t.Fatalf("Close: %v", err)
|
||||
}
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() { wg.Wait(); close(done) }()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("readers did not wake")
|
||||
}
|
||||
|
||||
for i, err := range errs {
|
||||
if !errors.Is(err, os.ErrClosed) {
|
||||
t.Errorf("reader %d: expected os.ErrClosed, got %v", i, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestPoll_ConcurrentWrite_NoRace hammers a single Poll queue from two writer
|
||||
// goroutines while a reader drains the other end of the pipe. The writers
|
||||
// overflow the pipe buffer, so both repeatedly park in blockOnWrite at the same
|
||||
// time — the exact scenario that raced on the old shared writePoll member
|
||||
// array. Run under -race; a shared-array regression trips the detector here.
|
||||
func TestPoll_ConcurrentWrite_NoRace(t *testing.T) {
|
||||
var fds [2]int
|
||||
require.NoError(t, unix.Pipe2(fds[:], unix.O_CLOEXEC))
|
||||
readFd, writeFd := fds[0], fds[1]
|
||||
|
||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = unix.Close(shutdownFd) })
|
||||
|
||||
p, err := newPoll(writeFd, shutdownFd)
|
||||
require.NoError(t, err)
|
||||
|
||||
const writers = 2
|
||||
const perWriter = 4000
|
||||
payload := make([]byte, 100)
|
||||
total := writers * perWriter * len(payload)
|
||||
|
||||
// Reader: drain the read end (blocking) until every writer's bytes are
|
||||
// consumed, so the writers keep making progress rather than wedging on a
|
||||
// permanently full pipe.
|
||||
readDone := make(chan struct{})
|
||||
go func() {
|
||||
defer close(readDone)
|
||||
buf := make([]byte, 4096)
|
||||
got := 0
|
||||
for got < total {
|
||||
n, rerr := unix.Read(readFd, buf)
|
||||
got += n
|
||||
if rerr != nil {
|
||||
if rerr == unix.EINTR {
|
||||
continue
|
||||
}
|
||||
return
|
||||
}
|
||||
if n == 0 { // EOF
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for w := 0; w < writers; w++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for i := 0; i < perWriter; i++ {
|
||||
if _, werr := p.Write(payload); werr != nil {
|
||||
t.Errorf("write: %v", werr)
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
select {
|
||||
case <-readDone:
|
||||
case <-time.After(10 * time.Second):
|
||||
t.Fatal("reader did not drain")
|
||||
}
|
||||
|
||||
require.NoError(t, p.Close())
|
||||
_ = unix.Close(readFd)
|
||||
}
|
||||
|
||||
// TestPoll_NewPoll_DoesNotCloseFdOnFailure pins the ownership rule: when
|
||||
// newPoll fails, it must leave fd open so the caller (pollQueueSet.Add's
|
||||
// callers in tun_linux.go) is the sole closer. If newPoll also closed fd,
|
||||
// the poll path would double-close on Add error. We force the failure with
|
||||
// an O_PATH descriptor: fcntl(F_SETFL) — which SetNonblock performs — is not
|
||||
// permitted on O_PATH fds and fails with EBADF, while the fd itself stays
|
||||
// open so we can observe that newPoll left it alone.
|
||||
func TestPoll_NewPoll_DoesNotCloseFdOnFailure(t *testing.T) {
|
||||
fd, err := unix.Open("/", unix.O_PATH|unix.O_CLOEXEC, 0)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = unix.Close(fd) })
|
||||
|
||||
p, err := newPoll(fd, 1)
|
||||
require.Error(t, err, "SetNonblock on an O_PATH fd should fail")
|
||||
require.Nil(t, p)
|
||||
|
||||
// If newPoll had closed fd, F_GETFD would report it closed. It staying
|
||||
// open proves newPoll left the fd for the caller to close exactly once.
|
||||
require.True(t, fdOpen(t, fd), "newPoll must not close fd on failure; caller is the sole closer")
|
||||
}
|
||||
|
||||
func TestPoll_Close_Idempotent(t *testing.T) {
|
||||
tf, err := newPoll(newReadPipe(t), 1)
|
||||
require.NoError(t, err)
|
||||
if err := tf.Close(); err != nil {
|
||||
t.Fatalf("first Close: %v", err)
|
||||
}
|
||||
if err := tf.Close(); err != nil {
|
||||
t.Fatalf("second Close should be a no-op, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// fdOpen reports whether fd currently refers to an open file description.
|
||||
// A closed (or never-allocated) fd makes F_GETFD fail with EBADF.
|
||||
func fdOpen(t *testing.T, fd int) bool {
|
||||
t.Helper()
|
||||
_, err := unix.FcntlInt(uintptr(fd), unix.F_GETFD, 0)
|
||||
if err == nil {
|
||||
return true
|
||||
}
|
||||
if errors.Is(err, unix.EBADF) {
|
||||
return false
|
||||
}
|
||||
t.Fatalf("unexpected fcntl(F_GETFD) error on fd %d: %v", fd, err)
|
||||
return false
|
||||
}
|
||||
|
||||
// TestPollQueueSet_Close_ClosesShutdownFd is the regression test for the
|
||||
// leaked shutdown eventfd: the container that owns shutdownFd must close it in
|
||||
// Close, and a second Close must be a safe no-op.
|
||||
func TestPollQueueSet_Close_ClosesShutdownFd(t *testing.T) {
|
||||
qs, err := NewPollQueueSet()
|
||||
require.NoError(t, err)
|
||||
c, ok := qs.(*pollQueueSet)
|
||||
require.True(t, ok)
|
||||
require.NoError(t, qs.Add(newReadPipe(t)))
|
||||
|
||||
shutdownFd := c.shutdownFd
|
||||
require.True(t, fdOpen(t, shutdownFd), "shutdown eventfd should be open before Close")
|
||||
|
||||
require.NoError(t, qs.Close())
|
||||
require.False(t, fdOpen(t, shutdownFd), "shutdown eventfd should be closed after Close")
|
||||
|
||||
// Second Close must not touch fds (shutdownFd is now -1) and must return nil.
|
||||
require.NoError(t, qs.Close())
|
||||
}
|
||||
|
||||
// TestOffloadQueueSet_Close_ClosesShutdownFd mirrors the poll regression test
|
||||
// for the GSO/offload queueset.
|
||||
func TestOffloadQueueSet_Close_ClosesShutdownFd(t *testing.T) {
|
||||
qs, err := NewOffloadQueueSet(false, slog.New(slog.DiscardHandler))
|
||||
require.NoError(t, err)
|
||||
c, ok := qs.(*offloadQueueSet)
|
||||
require.True(t, ok)
|
||||
require.NoError(t, qs.Add(newReadPipe(t)))
|
||||
|
||||
shutdownFd := c.shutdownFd
|
||||
require.True(t, fdOpen(t, shutdownFd), "shutdown eventfd should be open before Close")
|
||||
|
||||
require.NoError(t, qs.Close())
|
||||
require.False(t, fdOpen(t, shutdownFd), "shutdown eventfd should be closed after Close")
|
||||
|
||||
// Second Close must not touch fds (shutdownFd is now -1) and must return nil.
|
||||
require.NoError(t, qs.Close())
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package tio
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/tio/virtio"
|
||||
)
|
||||
|
||||
// protoFromGSOType maps a virtio_net_hdr gsoType to the GSOProto value the
|
||||
// segment-time helpers use. Returns an error for GSO_NONE or any unknown
|
||||
// value. The caller should only invoke this on a confirmed superpacket.
|
||||
func protoFromGSOType(t uint8) (GSOProto, error) {
|
||||
switch t &^ unix.VIRTIO_NET_HDR_GSO_ECN {
|
||||
case unix.VIRTIO_NET_HDR_GSO_TCPV4, unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
||||
return GSOProtoTCP, nil
|
||||
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
|
||||
return GSOProtoUDP, nil
|
||||
default:
|
||||
return 0, fmt.Errorf("unsupported virtio gso type: %d", t)
|
||||
}
|
||||
}
|
||||
|
||||
// gsoTypeFromProto is the reverse of protoFromGSOType
|
||||
func gsoTypeFromProto(proto GSOProto, ipVer uint8) uint8 {
|
||||
switch {
|
||||
case proto == GSOProtoUDP && (ipVer == 4 || ipVer == 6):
|
||||
return unix.VIRTIO_NET_HDR_GSO_UDP_L4
|
||||
case ipVer == 6:
|
||||
return unix.VIRTIO_NET_HDR_GSO_TCPV6
|
||||
case ipVer == 4:
|
||||
return unix.VIRTIO_NET_HDR_GSO_TCPV4
|
||||
default:
|
||||
return unix.VIRTIO_NET_HDR_GSO_NONE
|
||||
}
|
||||
}
|
||||
|
||||
// SegmentSuperpacket invokes fn once per segment of pkt.
|
||||
// For non-GSO pkts fn is called once with pkt.Bytes.
|
||||
// For GSO/USO superpackets, fn is called once per segment with a slice of pkt.Bytes holding that segment's plaintext
|
||||
// (a freshly-patched L3+L4 header sliced in front of the original payload chunk).
|
||||
// This slicing is destructive: pkt is consumed by this call.
|
||||
// Aborts and returns the first error from fn or from per-segment construction.
|
||||
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
||||
if !pkt.GSO.IsSuperpacket() {
|
||||
return fn(pkt.Bytes)
|
||||
}
|
||||
switch pkt.GSO.Proto {
|
||||
case GSOProtoTCP:
|
||||
return virtio.SegmentTCP(pkt.Bytes, pkt.GSO.HdrLen, pkt.GSO.CsumStart, pkt.GSO.Size, fn)
|
||||
case GSOProtoUDP:
|
||||
return virtio.SegmentUDP(pkt.Bytes, pkt.GSO.HdrLen, pkt.GSO.CsumStart, pkt.GSO.Size, fn)
|
||||
default:
|
||||
return fmt.Errorf("unsupported gso proto: %d", pkt.GSO.Proto)
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,75 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package virtio
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// Size is the on-wire length of struct virtio_net_hdr the kernel
|
||||
// prepends/expects on a TUN opened with IFF_VNET_HDR (TUNSETVNETHDRSZ
|
||||
// not set).
|
||||
const Size = 10
|
||||
|
||||
// Hdr is the Go view of the legacy virtio_net_hdr.
|
||||
type Hdr struct {
|
||||
Flags uint8
|
||||
gsoType uint8 //private to avoid mistakes wrt the 0x80 VIRTIO_NET_HDR_GSO_ECN flag, ORed with the other "GSO types"
|
||||
HdrLen uint16
|
||||
GSOSize uint16
|
||||
CsumStart uint16
|
||||
CsumOffset uint16
|
||||
}
|
||||
|
||||
func NewHeader(flags, gsoType uint8, hdrLen, gsoSize, csumStart, csumOffset uint16) Hdr {
|
||||
return Hdr{
|
||||
Flags: flags,
|
||||
gsoType: gsoType,
|
||||
HdrLen: hdrLen,
|
||||
GSOSize: gsoSize,
|
||||
CsumStart: csumStart,
|
||||
CsumOffset: csumOffset,
|
||||
}
|
||||
}
|
||||
|
||||
// Decode reads a virtio_net_hdr in host byte order (TUN default; we never
|
||||
// call TUNSETVNETLE so the kernel matches our endianness).
|
||||
func (h *Hdr) Decode(b []byte) {
|
||||
h.Flags = b[0]
|
||||
h.gsoType = b[1]
|
||||
h.HdrLen = binary.NativeEndian.Uint16(b[2:4])
|
||||
h.GSOSize = binary.NativeEndian.Uint16(b[4:6])
|
||||
h.CsumStart = binary.NativeEndian.Uint16(b[6:8])
|
||||
h.CsumOffset = binary.NativeEndian.Uint16(b[8:10])
|
||||
}
|
||||
|
||||
func EncodeHeader(b []byte, flags, gsoType uint8, hdrLen, gsoSize, csumStart, csumOffset uint16) {
|
||||
b[0] = flags
|
||||
b[1] = gsoType
|
||||
binary.NativeEndian.PutUint16(b[2:4], hdrLen)
|
||||
binary.NativeEndian.PutUint16(b[4:6], gsoSize)
|
||||
binary.NativeEndian.PutUint16(b[6:8], csumStart)
|
||||
binary.NativeEndian.PutUint16(b[8:10], csumOffset)
|
||||
}
|
||||
|
||||
// Encode is the inverse of Decode: writes the virtio_net_hdr fields into b
|
||||
// (must be at least Size bytes). Used to emit a TSO superpacket on egress.
|
||||
func (h *Hdr) Encode(b []byte) {
|
||||
EncodeHeader(b, h.Flags, h.gsoType, h.HdrLen, h.GSOSize, h.CsumStart, h.CsumOffset)
|
||||
}
|
||||
|
||||
// GSOType returns gsoType with the ECN-flag masked out
|
||||
func (h *Hdr) GSOType() uint8 {
|
||||
return h.gsoType &^ unix.VIRTIO_NET_HDR_GSO_ECN
|
||||
}
|
||||
|
||||
func (h *Hdr) HasECNFlag() bool {
|
||||
return h.gsoType&unix.VIRTIO_NET_HDR_GSO_ECN != 0
|
||||
}
|
||||
|
||||
func (h *Hdr) SetGSOType(x uint8) {
|
||||
h.gsoType = x
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
//go:build !linux || android
|
||||
|
||||
package virtio
|
||||
@@ -0,0 +1,441 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
// Package virtio implements the pure validation, header-correction, and
|
||||
// per-segment slicing logic for kernel-supplied TSO/USO superpackets on
|
||||
// IFF_VNET_HDR TUN devices. It is FD-free and depends only on the byte
|
||||
// layout of the virtio_net_hdr and the IP/TCP/UDP headers it describes,
|
||||
// so it can be unit-tested in isolation from the tio Queue runtime.
|
||||
package virtio
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/checksum"
|
||||
)
|
||||
|
||||
// Protocol header size bounds used to validate / cap kernel-supplied offsets.
|
||||
const (
|
||||
ipv4HeaderMinLen = 20 // IHL=5, no options
|
||||
ipv4HeaderMaxLen = 60 // IHL=15, max options
|
||||
ipv6FixedLen = 40 // IPv6 base header; extensions would extend this
|
||||
tcpHeaderMinLen = 20 // data-offset=5, no options
|
||||
tcpHeaderMaxLen = 60 // data-offset=15, max options
|
||||
)
|
||||
|
||||
// maxSegHdrLen bounds the L3+L4 header we snapshot before stamping each segment.
|
||||
// The largest header the segmenter supports is IPv4 (max IHL 60) plus TCP (max data-offset 60) = 120 bytes
|
||||
const maxSegHdrLen = ipv4HeaderMaxLen + tcpHeaderMaxLen // 120
|
||||
|
||||
// Byte offsets inside an IPv4 header.
|
||||
const (
|
||||
ipv4TotalLenOff = 2
|
||||
ipv4IDOff = 4
|
||||
ipv4ChecksumOff = 10
|
||||
ipv4SrcOff = 12
|
||||
ipv4AddrsEnd = 20 // end of dst address (ipv4SrcOff + 2*4)
|
||||
)
|
||||
|
||||
// Byte offsets inside an IPv6 header.
|
||||
const (
|
||||
ipv6PayloadLenOff = 4
|
||||
ipv6SrcOff = 8
|
||||
ipv6AddrsEnd = 40 // end of dst address (ipv6SrcOff + 2*16)
|
||||
)
|
||||
|
||||
// Byte offsets inside a TCP header (relative to its start, i.e. csumStart).
|
||||
const (
|
||||
tcpSeqOff = 4
|
||||
tcpDataOffOff = 12 // upper nibble is header len in 32-bit words
|
||||
tcpFlagsOff = 13
|
||||
tcpChecksumOff = 16
|
||||
)
|
||||
|
||||
// UDP header is fixed at 8 bytes: {sport, dport, length, checksum}.
|
||||
const (
|
||||
udpHeaderLen = 8
|
||||
udpLengthOff = 4
|
||||
udpChecksumOff = 6
|
||||
)
|
||||
|
||||
var errPacketTooShort = errors.New("packet too short")
|
||||
|
||||
// tcpFinPshMask is cleared on every segment except the last of a TSO burst.
|
||||
const tcpFinPshMask = 0x09 // FIN(0x01) | PSH(0x08)
|
||||
|
||||
// tcpCwrFlag is cleared on every segment except the first.
|
||||
// Per RFC 3168 §6.1.2 the CWR bit signals a one-shot transition (the sender just halved its window)
|
||||
// and must appear on the first segment of a TSO burst only.
|
||||
const tcpCwrFlag = 0x80
|
||||
|
||||
// CheckValid rejects packets whose virtio_net_hdr/IP combination would
|
||||
// cause a downstream miscompute. The TUN should never emit RSC_INFO and
|
||||
// the GSO type must agree with the IP version nibble.
|
||||
func CheckValid(pkt []byte, hdr Hdr) error {
|
||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_RSC_INFO != 0 {
|
||||
return fmt.Errorf("virtio RSC_INFO flag not supported on TUN reads")
|
||||
}
|
||||
if len(pkt) < ipv4HeaderMinLen {
|
||||
return errPacketTooShort
|
||||
}
|
||||
ipVersion := pkt[0] >> 4
|
||||
if ipVersion == 6 && len(pkt) < ipv6FixedLen {
|
||||
return errPacketTooShort
|
||||
}
|
||||
|
||||
gsoType := hdr.GSOType()
|
||||
if gsoType != unix.VIRTIO_NET_HDR_GSO_NONE && hdr.GSOSize == 0 {
|
||||
// A GSO type with no segment size would dodge IsSuperpacket() downstream and
|
||||
// travel as a plain jumbo datagram with an unfinished checksum.
|
||||
return fmt.Errorf("virtio GSO type %#x with zero gso_size", hdr.gsoType)
|
||||
}
|
||||
if hdr.HasECNFlag() && !(gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV4 || gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV6) {
|
||||
return fmt.Errorf("virtio GSO_ECN qualifier on non-TCP GSO type %#x", hdr.gsoType)
|
||||
}
|
||||
switch gsoType {
|
||||
case unix.VIRTIO_NET_HDR_GSO_TCPV4:
|
||||
if ipVersion != 4 {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
}
|
||||
case unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
||||
if ipVersion != 6 {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
}
|
||||
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
|
||||
// USO carries either v4 or v6; the leading nibble disambiguates.
|
||||
if !(ipVersion == 4 || ipVersion == 6) {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
}
|
||||
default:
|
||||
if !(ipVersion == 6 || ipVersion == 4) {
|
||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.gsoType)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// CorrectHdrLen rewrites hdr.HdrLen based on the actual transport header length read out of pkt.
|
||||
// The kernel's hdr.HdrLen on the FORWARD path can be the length of the entire first packet, so we don't trust it.
|
||||
func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
|
||||
// Thank you wireguard-go for documenting these edge-cases
|
||||
// Don't trust hdr.hdrLen from the kernel as it can be equal to the length
|
||||
// of the entire first packet when the kernel is handling it as part of a FORWARD path.
|
||||
// Instead, parse the transport header length and add it onto csumStart, which is synonymous for IP header length.
|
||||
|
||||
if hdr.GSOType() == unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
|
||||
hdr.HdrLen = hdr.CsumStart + 8
|
||||
} else {
|
||||
if len(pkt) <= int(hdr.CsumStart+tcpDataOffOff) {
|
||||
return errors.New("packet is too short")
|
||||
}
|
||||
|
||||
tcpHLen := uint16(pkt[hdr.CsumStart+tcpDataOffOff] >> 4 * 4)
|
||||
if tcpHLen < tcpHeaderMinLen || tcpHLen > tcpHeaderMaxLen {
|
||||
return fmt.Errorf("tcp header len is invalid: %d", tcpHLen)
|
||||
}
|
||||
hdr.HdrLen = hdr.CsumStart + tcpHLen
|
||||
}
|
||||
|
||||
if len(pkt) < int(hdr.HdrLen) {
|
||||
return fmt.Errorf("length of packet (%d) < virtioNetHdr.HdrLen (%d)", len(pkt), hdr.HdrLen)
|
||||
}
|
||||
|
||||
if hdr.HdrLen < hdr.CsumStart {
|
||||
return fmt.Errorf("virtioNetHdr.HdrLen (%d) < virtioNetHdr.CsumStart (%d)", hdr.HdrLen, hdr.CsumStart)
|
||||
}
|
||||
cSumAt := int(hdr.CsumStart + hdr.CsumOffset)
|
||||
if cSumAt+1 >= len(pkt) {
|
||||
return fmt.Errorf("end of checksum offset (%d) exceeds packet length (%d)", cSumAt+1, len(pkt))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// segCount returns how many segments a payload of payLen bytes splits into at gsoSize,
|
||||
// with a floor of one so a header-only superpacket still yields a single segment.
|
||||
func segCount(payLen, gsoSize int) int {
|
||||
n := (payLen + gsoSize - 1) / gsoSize
|
||||
if n == 0 {
|
||||
return 1
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
// basePseudoSum folds the part of the L4 pseudo-header sum that is identical
|
||||
// for every segment: the source and destination addresses plus the protocol
|
||||
// number. The per-segment L4 length is added by the caller inside the loop.
|
||||
func basePseudoSum(pkt []byte, isV4 bool, proto uint32) uint32 {
|
||||
if isV4 {
|
||||
return uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0)) + proto
|
||||
}
|
||||
return uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0)) + proto
|
||||
}
|
||||
|
||||
// baseIPv4HdrSum folds the IPv4 header checksum over the fields that stay constant across segments.
|
||||
// csumStart is the L3 header length, which bounds a valid IHL.
|
||||
func baseIPv4HdrSum(pkt []byte, csumStart int) (uint32, error) {
|
||||
ihl := int(pkt[0]&0x0f) * 4
|
||||
if ihl < ipv4HeaderMinLen || ihl > csumStart {
|
||||
return 0, fmt.Errorf("bad IPv4 IHL: %d", ihl)
|
||||
}
|
||||
// total_len, the ID, and the checksum field itself are excluded: all three are rewritten per segment.
|
||||
sum := uint32(checksum.Checksum(pkt[:ihl], 0))
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4TotalLenOff : ipv4TotalLenOff+2]))
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4ChecksumOff : ipv4ChecksumOff+2]))
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2]))
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
return sum, nil
|
||||
}
|
||||
|
||||
// baseTCPHdrSum folds the TCP header checksum over everything the segment loop does not rewrite
|
||||
func baseTCPHdrSum(pkt []byte, csumStart, headerLen int) uint32 {
|
||||
seq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
|
||||
flags := uint16(pkt[csumStart+tcpFlagsOff])
|
||||
|
||||
sum := uint32(checksum.Checksum(pkt[csumStart:headerLen], 0))
|
||||
sum += uint32(^uint16(seq >> 16))
|
||||
sum += uint32(^uint16(seq))
|
||||
sum += uint32(^flags)
|
||||
sum += uint32(^binary.BigEndian.Uint16(pkt[csumStart+tcpChecksumOff : csumStart+tcpChecksumOff+2]))
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
return sum
|
||||
}
|
||||
|
||||
// SegmentTCP walks a TSO superpacket pkt, yielding each segment as a slice into pkt.
|
||||
// Per-segment plaintext is laid out by stamping a copy of the original L3+L4 header into pkt at offset i*gsoSize,
|
||||
// where it sits immediately before that segment's payload chunk in the original buffer.
|
||||
// pkt is consumed by this call and must not be inspected by the caller after the final yield.
|
||||
func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
|
||||
if gsoSizeU == 0 {
|
||||
return fmt.Errorf("gso_size is zero")
|
||||
}
|
||||
if csumStartU == 0 {
|
||||
return fmt.Errorf("csum_start is zero")
|
||||
}
|
||||
|
||||
headerLen := int(hdrLenU)
|
||||
csumStart := int(csumStartU)
|
||||
if headerLen > maxSegHdrLen {
|
||||
return fmt.Errorf("header len %d exceeds max %d", headerLen, maxSegHdrLen)
|
||||
}
|
||||
isV4 := pkt[0]>>4 == 4
|
||||
|
||||
tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4
|
||||
payLen := len(pkt) - headerLen
|
||||
gsoSize := int(gsoSizeU)
|
||||
numSeg := segCount(payLen, gsoSize)
|
||||
|
||||
origSeq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
|
||||
origFlags := pkt[csumStart+tcpFlagsOff]
|
||||
|
||||
baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_TCP)
|
||||
baseTcpHdrSum := baseTCPHdrSum(pkt, csumStart, headerLen)
|
||||
|
||||
var origIPID uint16
|
||||
var baseIPHdrSum uint32
|
||||
if isV4 {
|
||||
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
|
||||
var err error
|
||||
// TSO bumps the ID per segment, so it stays out of the base sum.
|
||||
baseIPHdrSum, err = baseIPv4HdrSum(pkt, csumStart)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// Snapshot the pristine L3+L4 header once. '
|
||||
// Every segment's header is stamped from this copy, so overlapping stamps (gsoSize < headerLen) can never corrupt the source.
|
||||
var savedHdr [maxSegHdrLen]byte
|
||||
copy(savedHdr[:headerLen], pkt[:headerLen])
|
||||
|
||||
for i := 0; i < numSeg; i++ {
|
||||
segStart := i * gsoSize
|
||||
segEnd := segStart + gsoSize
|
||||
if segEnd > payLen {
|
||||
segEnd = payLen
|
||||
}
|
||||
segPayLen := segEnd - segStart
|
||||
segLen := headerLen + segPayLen
|
||||
headerOff := i * gsoSize
|
||||
|
||||
// Stamp the header into place immediately before this segment's payload, sourced from the snapshot.
|
||||
// The per-segment patches below overwrite the variable fields. (seq/flags/cksum/totalLen/id)
|
||||
if i > 0 {
|
||||
// Iter 0's header is already at pkt[:headerLen] (identical to savedHdr), so only i >= 1 needs the stamp
|
||||
copy(pkt[headerOff:headerOff+headerLen], savedHdr[:headerLen])
|
||||
}
|
||||
seg := pkt[headerOff : headerOff+segLen]
|
||||
|
||||
segSeq := origSeq + uint32(segStart)
|
||||
segFlags := origFlags
|
||||
if i != 0 {
|
||||
segFlags &^= tcpCwrFlag
|
||||
}
|
||||
if i != numSeg-1 {
|
||||
segFlags &^= tcpFinPshMask
|
||||
}
|
||||
totalLen := segLen
|
||||
|
||||
if isV4 {
|
||||
segID := origIPID + uint16(i)
|
||||
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
|
||||
binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID)
|
||||
ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID)
|
||||
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
|
||||
} else {
|
||||
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
|
||||
}
|
||||
|
||||
binary.BigEndian.PutUint32(seg[csumStart+tcpSeqOff:csumStart+tcpSeqOff+4], segSeq)
|
||||
seg[csumStart+tcpFlagsOff] = segFlags
|
||||
|
||||
tcpLen := tcpHdrLen + segPayLen
|
||||
// Payload bytes still live at their original offset in pkt.
|
||||
// The header slide above only writes into pkt[i*GSOSize : i*GSOSize+header], which is the tail of seg_{i-1}'s payload (already consumed)
|
||||
// and never overlaps seg_i's own payload at pkt[header+i*GSOSize : header+(i+1)*GSOSize].
|
||||
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
|
||||
wide := uint64(baseTcpHdrSum) + uint64(paySum) + uint64(baseProtoSum)
|
||||
wide += uint64(segSeq) + uint64(segFlags) + uint64(tcpLen)
|
||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
||||
binary.BigEndian.PutUint16(seg[csumStart+tcpChecksumOff:csumStart+tcpChecksumOff+2], foldComplement(uint32(wide)))
|
||||
|
||||
if err := yield(seg); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// SegmentUDP walks a USO superpacket, stamping a per-segment-patched copy of the original L3+L4 header
|
||||
// into pkt at offset i*GSOSize and yielding pkt[i*GSOSize:i*GSOSize+segLen] to the caller.
|
||||
// Per-segment patches are total_len + IPv4 csum (or IPv6 payload_len) plus the UDP length and checksum.
|
||||
// pkt is consumed destructively.
|
||||
func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
|
||||
if gsoSizeU == 0 {
|
||||
return fmt.Errorf("gso_size is zero")
|
||||
}
|
||||
if csumStartU == 0 {
|
||||
return fmt.Errorf("csum_start is zero")
|
||||
}
|
||||
|
||||
isV4 := pkt[0]>>4 == 4
|
||||
headerLen := int(hdrLenU)
|
||||
csumStart := int(csumStartU)
|
||||
if headerLen > maxSegHdrLen {
|
||||
return fmt.Errorf("header len %d exceeds max %d", headerLen, maxSegHdrLen)
|
||||
}
|
||||
if headerLen-csumStart != udpHeaderLen {
|
||||
return fmt.Errorf("udp header len mismatch: %d", headerLen-csumStart)
|
||||
}
|
||||
|
||||
payLen := len(pkt) - headerLen
|
||||
gsoSize := int(gsoSizeU)
|
||||
numSeg := segCount(payLen, gsoSize)
|
||||
|
||||
baseProtoSum := basePseudoSum(pkt, isV4, unix.IPPROTO_UDP)
|
||||
|
||||
var origIPID uint16
|
||||
var baseIPHdrSum uint32
|
||||
if isV4 {
|
||||
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
|
||||
var err error
|
||||
// Software UDP GSO bumps the ID per segment just like TSO
|
||||
// (inet_gso_segment's fixed-ID case is TCP-only), so it stays out of the base sum.
|
||||
baseIPHdrSum, err = baseIPv4HdrSum(pkt, csumStart)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// Snapshot the pristine L3+L4 header once and stamp every segment from it
|
||||
var savedHdr [maxSegHdrLen]byte
|
||||
copy(savedHdr[:headerLen], pkt[:headerLen])
|
||||
|
||||
for i := 0; i < numSeg; i++ {
|
||||
segStart := i * gsoSize
|
||||
segEnd := segStart + gsoSize
|
||||
if segEnd > payLen {
|
||||
segEnd = payLen
|
||||
}
|
||||
segPayLen := segEnd - segStart
|
||||
segLen := headerLen + segPayLen
|
||||
headerOff := i * gsoSize
|
||||
|
||||
if i > 0 {
|
||||
copy(pkt[headerOff:headerOff+headerLen], savedHdr[:headerLen])
|
||||
}
|
||||
seg := pkt[headerOff : headerOff+segLen]
|
||||
|
||||
totalLen := segLen
|
||||
udpLen := udpHeaderLen + segPayLen
|
||||
|
||||
if isV4 {
|
||||
segID := origIPID + uint16(i)
|
||||
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
|
||||
binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID)
|
||||
ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID)
|
||||
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
|
||||
} else {
|
||||
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
|
||||
}
|
||||
|
||||
binary.BigEndian.PutUint16(seg[csumStart+udpLengthOff:csumStart+udpLengthOff+2], uint16(udpLen))
|
||||
|
||||
// Sum the UDP header (length just written, checksum zeroed) together with
|
||||
// this segment's payload in one pass, seeded with the pseudo-header sum.
|
||||
seg[csumStart+udpChecksumOff], seg[csumStart+udpChecksumOff+1] = 0, 0
|
||||
pseudo := baseProtoSum + uint32(udpLen)
|
||||
pseudo = (pseudo & 0xffff) + (pseudo >> 16)
|
||||
pseudo = (pseudo & 0xffff) + (pseudo >> 16)
|
||||
csum := ^checksum.Checksum(seg[csumStart:], uint16(pseudo))
|
||||
if csum == 0 {
|
||||
csum = 0xffff
|
||||
}
|
||||
binary.BigEndian.PutUint16(seg[csumStart+udpChecksumOff:csumStart+udpChecksumOff+2], csum)
|
||||
|
||||
if err := yield(seg); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// FinishChecksum computes the L4 checksum for a non-GSO packet that the kernel handed us with NEEDS_CSUM set.
|
||||
// CsumStart / CsumOffset point at the 16-bit checksum field.
|
||||
// We zero it, fold a full sum from the partial one that the kernel provided, and store the result.
|
||||
func FinishChecksum(seg []byte, hdr Hdr) error {
|
||||
cs := int(hdr.CsumStart)
|
||||
co := int(hdr.CsumOffset)
|
||||
if cs+co+2 > len(seg) {
|
||||
return fmt.Errorf("csum offsets out of range: start=%d offset=%d len=%d", cs, co, len(seg))
|
||||
}
|
||||
// The kernel stores a partial pseudo-header sum at [cs+co:]; sum over the
|
||||
// L4 region starting at cs, folding the prior partial in as the seed.
|
||||
partial := binary.BigEndian.Uint16(seg[cs+co : cs+co+2])
|
||||
seg[cs+co] = 0
|
||||
seg[cs+co+1] = 0
|
||||
csum := ^checksum.Checksum(seg[cs:], partial)
|
||||
// RFC 768: UDP transmits a computed zero as all ones, since all-zero is the reserved "no checksum" value.
|
||||
if co == udpChecksumOff && csum == 0 {
|
||||
csum = 0xffff
|
||||
}
|
||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], csum)
|
||||
return nil
|
||||
}
|
||||
|
||||
// foldComplement folds a 32-bit one's-complement partial sum to 16 bits and
|
||||
// complements it, yielding the on-wire Internet checksum value.
|
||||
func foldComplement(sum uint32) uint16 {
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
sum = (sum & 0xffff) + (sum >> 16)
|
||||
return ^uint16(sum)
|
||||
}
|
||||
@@ -0,0 +1,602 @@
|
||||
//go:build linux && !android
|
||||
// +build linux,!android
|
||||
|
||||
package virtio
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"testing"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/checksum"
|
||||
)
|
||||
|
||||
// verifyChecksum confirms that the one's-complement sum across b, seeded with
|
||||
// a folded pseudo-header sum, equals all-ones (a valid on-wire checksum).
|
||||
// A corrupted header stamped into a segment makes this fail even when the
|
||||
// checksum field itself was computed from the (pristine) base sums, because
|
||||
// the bytes the receiver would sum no longer match what was checksummed.
|
||||
func verifyChecksum(b []byte, pseudo uint16) bool {
|
||||
return checksum.Checksum(b, pseudo) == 0xffff
|
||||
}
|
||||
|
||||
// pseudoHeaderIPv4 folds the TCP/UDP pseudo-header sum from a segment's own
|
||||
// address and length fields, used to independently verify its L4 checksum.
|
||||
func pseudoHeaderIPv4(src, dst []byte, proto byte, l4Len int) uint16 {
|
||||
s := uint32(checksum.Checksum(src, 0)) + uint32(checksum.Checksum(dst, 0))
|
||||
s += uint32(proto) + uint32(l4Len)
|
||||
s = (s & 0xffff) + (s >> 16)
|
||||
s = (s & 0xffff) + (s >> 16)
|
||||
return uint16(s)
|
||||
}
|
||||
|
||||
// buildTCPv4Super constructs a synthetic IPv4/TCP TSO superpacket with a
|
||||
// payload of payLen bytes and returns it alongside the header fields the
|
||||
// segmenter needs. The header is a fixed 40 bytes (20 IPv4 + 20 TCP).
|
||||
func buildTCPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
|
||||
const ipLen = 20
|
||||
const tcpLen = 20
|
||||
pkt = make([]byte, ipLen+tcpLen+payLen)
|
||||
|
||||
// IPv4 header.
|
||||
pkt[0] = 0x45 // version 4, IHL 5
|
||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+tcpLen+payLen))
|
||||
binary.BigEndian.PutUint16(pkt[4:6], 0x4242) // ID
|
||||
pkt[8] = 64 // TTL
|
||||
pkt[9] = unix.IPPROTO_TCP
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1}) // src
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2}) // dst
|
||||
|
||||
// TCP header.
|
||||
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
|
||||
binary.BigEndian.PutUint16(pkt[22:24], 80) // dport
|
||||
binary.BigEndian.PutUint32(pkt[24:28], 10000) // seq
|
||||
binary.BigEndian.PutUint32(pkt[28:32], 20000) // ack
|
||||
pkt[32] = 0x50 // data offset 5 words
|
||||
pkt[33] = 0x18 // ACK | PSH
|
||||
binary.BigEndian.PutUint16(pkt[34:36], 65535) // window
|
||||
|
||||
for i := 0; i < payLen; i++ {
|
||||
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
|
||||
}
|
||||
return pkt, ipLen + tcpLen, ipLen
|
||||
}
|
||||
|
||||
// buildUDPv4Super constructs a synthetic IPv4/UDP USO superpacket with a
|
||||
// payload of payLen bytes. Header is a fixed 28 bytes (20 IPv4 + 8 UDP).
|
||||
func buildUDPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
|
||||
const ipLen = 20
|
||||
const udpLen = 8
|
||||
pkt = make([]byte, ipLen+udpLen+payLen)
|
||||
|
||||
pkt[0] = 0x45
|
||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+udpLen+payLen))
|
||||
binary.BigEndian.PutUint16(pkt[4:6], 0x4242)
|
||||
pkt[8] = 64
|
||||
pkt[9] = unix.IPPROTO_UDP
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||
|
||||
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
|
||||
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
|
||||
|
||||
for i := 0; i < payLen; i++ {
|
||||
pkt[ipLen+udpLen+i] = byte(i & 0xff)
|
||||
}
|
||||
return pkt, ipLen + udpLen, ipLen
|
||||
}
|
||||
|
||||
// collectTCP segments a fresh copy of pkt and returns each segment as an
|
||||
// independent slice so assertions can run after segmentation completes.
|
||||
func collectTCP(t *testing.T, pkt []byte, hdrLen, csumStart, gsoSize uint16) [][]byte {
|
||||
t.Helper()
|
||||
work := append([]byte(nil), pkt...)
|
||||
var out [][]byte
|
||||
err := SegmentTCP(work, hdrLen, csumStart, gsoSize, func(seg []byte) error {
|
||||
out = append(out, append([]byte(nil), seg...))
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SegmentTCP: %v", err)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func collectUDP(t *testing.T, pkt []byte, hdrLen, csumStart, gsoSize uint16) [][]byte {
|
||||
t.Helper()
|
||||
work := append([]byte(nil), pkt...)
|
||||
var out [][]byte
|
||||
err := SegmentUDP(work, hdrLen, csumStart, gsoSize, func(seg []byte) error {
|
||||
out = append(out, append([]byte(nil), seg...))
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SegmentUDP: %v", err)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// TestSegmentTCPHeaderNotCorrupted is the regression test for the in-place
|
||||
// header-slide bug: when gsoSize < headerLen the old code stamped each
|
||||
// segment's header from pkt[:headerLen], which had already been overwritten
|
||||
// by the previous segment's overlapping stamp, so segments 2..n carried a
|
||||
// corrupted header (garbage src/dst/ports/seq). Every segment must instead
|
||||
// carry the ORIGINAL constant header fields with correct per-segment seq.
|
||||
func TestSegmentTCPHeaderNotCorrupted(t *testing.T) {
|
||||
const origSeq = 10000
|
||||
cases := []struct {
|
||||
name string
|
||||
payLen int
|
||||
gsoSize uint16
|
||||
}{
|
||||
// gsoSize (8) < headerLen (40): the bug's trigger. Even split.
|
||||
{"small-gso-even", 40, 8},
|
||||
// gsoSize (8) < headerLen (40) with a short final segment.
|
||||
{"small-gso-odd-tail", 44, 8},
|
||||
// gsoSize (100) >= headerLen (40): the normal path, must still work.
|
||||
{"normal-gso", 250, 100},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
pkt, hdrLen, csumStart := buildTCPv4Super(tc.payLen)
|
||||
gso := int(tc.gsoSize)
|
||||
wantSeg := (tc.payLen + gso - 1) / gso
|
||||
segs := collectTCP(t, pkt, hdrLen, csumStart, tc.gsoSize)
|
||||
if len(segs) != wantSeg {
|
||||
t.Fatalf("got %d segments, want %d", len(segs), wantSeg)
|
||||
}
|
||||
|
||||
off := 0
|
||||
for i, seg := range segs {
|
||||
// Constant header fields must be identical to the original in
|
||||
// EVERY segment. These are exactly the bytes the old code
|
||||
// corrupted in segments 2..n.
|
||||
if got := seg[0]; got != 0x45 {
|
||||
t.Errorf("seg %d: version/IHL byte=%#x want 0x45", i, got)
|
||||
}
|
||||
if seg[9] != unix.IPPROTO_TCP {
|
||||
t.Errorf("seg %d: proto=%d want %d", i, seg[9], unix.IPPROTO_TCP)
|
||||
}
|
||||
if !bytes.Equal(seg[12:16], []byte{10, 0, 0, 1}) {
|
||||
t.Errorf("seg %d: src=%v want [10 0 0 1]", i, seg[12:16])
|
||||
}
|
||||
if !bytes.Equal(seg[16:20], []byte{10, 0, 0, 2}) {
|
||||
t.Errorf("seg %d: dst=%v want [10 0 0 2]", i, seg[16:20])
|
||||
}
|
||||
if sport := binary.BigEndian.Uint16(seg[20:22]); sport != 12345 {
|
||||
t.Errorf("seg %d: sport=%d want 12345", i, sport)
|
||||
}
|
||||
if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 80 {
|
||||
t.Errorf("seg %d: dport=%d want 80", i, dport)
|
||||
}
|
||||
if ack := binary.BigEndian.Uint32(seg[28:32]); ack != 20000 {
|
||||
t.Errorf("seg %d: ack=%d want 20000", i, ack)
|
||||
}
|
||||
if seg[32] != 0x50 {
|
||||
t.Errorf("seg %d: data-offset byte=%#x want 0x50", i, seg[32])
|
||||
}
|
||||
|
||||
// Per-segment seq must advance by the payload offset.
|
||||
segStart := i * gso
|
||||
if seq := binary.BigEndian.Uint32(seg[24:28]); seq != uint32(origSeq+segStart) {
|
||||
t.Errorf("seg %d: seq=%d want %d", i, seq, origSeq+segStart)
|
||||
}
|
||||
|
||||
// Payload bytes must be the original contiguous slice.
|
||||
segPayLen := len(seg) - int(hdrLen)
|
||||
wantPay := make([]byte, segPayLen)
|
||||
for k := 0; k < segPayLen; k++ {
|
||||
wantPay[k] = byte((off + k) & 0xff)
|
||||
}
|
||||
if !bytes.Equal(seg[hdrLen:], wantPay) {
|
||||
t.Errorf("seg %d: payload mismatch", i)
|
||||
}
|
||||
off += segPayLen
|
||||
|
||||
// End-to-end: the stamped header must checksum-verify. A
|
||||
// corrupted header fails here because the written checksum was
|
||||
// derived from the pristine header.
|
||||
if !verifyChecksum(seg[:20], 0) {
|
||||
t.Errorf("seg %d: bad IPv4 header checksum", i)
|
||||
}
|
||||
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_TCP, len(seg)-20)
|
||||
if !verifyChecksum(seg[20:], psum) {
|
||||
t.Errorf("seg %d: bad TCP checksum", i)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestCorrectHdrLenChecksumBound guards the checksum-field bounds check in
|
||||
// CorrectHdrLen. The checksum field sits at CsumStart+CsumOffset, so the check
|
||||
// must be computed from CsumStart+CsumOffset — NOT CsumStart+CsumStart, a
|
||||
// regression that doubled CsumStart and thus over-tightened the bound (since
|
||||
// CsumOffset, 6 for UDP / 16 for TCP, is always < CsumStart >= 20). That bogus
|
||||
// bound spuriously rejected valid small USO superpackets in decodeRead.
|
||||
func TestCorrectHdrLenChecksumBound(t *testing.T) {
|
||||
// A valid IPv4 USO superpacket: 20B IPv4 + 8B UDP + two 6-byte segments
|
||||
// (payload 12) = 40 bytes total. CsumStart=20, CsumOffset=6, so the UDP
|
||||
// checksum field lives at bytes 26..27, comfortably inside the 40-byte
|
||||
// packet. The OLD formula computed cSumAt = CsumStart+CsumStart = 40 and
|
||||
// rejected on cSumAt+1 (41) >= len(pkt) (40); the fix (CsumStart+CsumOffset
|
||||
// = 26) accepts. This case FAILS against the CsumStart+CsumStart regression.
|
||||
t.Run("valid-small-uso-accepted", func(t *testing.T) {
|
||||
pkt, _, csumStart := buildUDPv4Super(12) // total len 40
|
||||
hdr := NewHeader(
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
|
||||
0, /*hdrLen*/
|
||||
6, /*gsoSize: two 6-byte segments*/
|
||||
csumStart, /*csumStart*/
|
||||
6, /*csumOffset*/
|
||||
)
|
||||
if err := CorrectHdrLen(pkt, &hdr); err != nil {
|
||||
t.Fatalf("CorrectHdrLen rejected a valid 40-byte USO superpacket: %v", err)
|
||||
}
|
||||
if hdr.HdrLen != csumStart+udpHeaderLen {
|
||||
t.Errorf("HdrLen = %d, want %d", hdr.HdrLen, csumStart+udpHeaderLen)
|
||||
}
|
||||
})
|
||||
|
||||
// A genuinely-too-short packet: CsumStart=20, CsumOffset=6 means the
|
||||
// checksum field would end at byte 27, but the packet is only 25 bytes
|
||||
// (CsumStart+CsumOffset+2 = 28 > 25). CorrectHdrLen must still reject it.
|
||||
t.Run("too-short-rejected", func(t *testing.T) {
|
||||
pkt := make([]byte, 25)
|
||||
pkt[0] = 0x45 // IPv4, IHL 5
|
||||
hdr := NewHeader(
|
||||
unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, /*flags*/
|
||||
unix.VIRTIO_NET_HDR_GSO_UDP_L4, /*gsoType*/
|
||||
0, /*hdrLen*/
|
||||
6, /*gsoSize*/
|
||||
20, /*csumStart*/
|
||||
6, /*csumOffset*/
|
||||
)
|
||||
if err := CorrectHdrLen(pkt, &hdr); err == nil {
|
||||
t.Fatalf("CorrectHdrLen accepted a too-short (25-byte) packet")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestSegmentUDPHeaderNotCorrupted is the USO counterpart: SegmentUDP performs
|
||||
// the same header stamp and must be correct when gsoSize < headerLen.
|
||||
func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
payLen int
|
||||
gsoSize uint16
|
||||
}{
|
||||
{"small-gso-even", 40, 8},
|
||||
{"small-gso-odd-tail", 44, 8},
|
||||
{"normal-gso", 250, 100},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
pkt, hdrLen, csumStart := buildUDPv4Super(tc.payLen)
|
||||
gso := int(tc.gsoSize)
|
||||
wantSeg := (tc.payLen + gso - 1) / gso
|
||||
segs := collectUDP(t, pkt, hdrLen, csumStart, tc.gsoSize)
|
||||
if len(segs) != wantSeg {
|
||||
t.Fatalf("got %d segments, want %d", len(segs), wantSeg)
|
||||
}
|
||||
|
||||
off := 0
|
||||
for i, seg := range segs {
|
||||
if got := seg[0]; got != 0x45 {
|
||||
t.Errorf("seg %d: version/IHL byte=%#x want 0x45", i, got)
|
||||
}
|
||||
if seg[9] != unix.IPPROTO_UDP {
|
||||
t.Errorf("seg %d: proto=%d want %d", i, seg[9], unix.IPPROTO_UDP)
|
||||
}
|
||||
if !bytes.Equal(seg[12:16], []byte{10, 0, 0, 1}) {
|
||||
t.Errorf("seg %d: src=%v want [10 0 0 1]", i, seg[12:16])
|
||||
}
|
||||
if !bytes.Equal(seg[16:20], []byte{10, 0, 0, 2}) {
|
||||
t.Errorf("seg %d: dst=%v want [10 0 0 2]", i, seg[16:20])
|
||||
}
|
||||
if sport := binary.BigEndian.Uint16(seg[20:22]); sport != 12345 {
|
||||
t.Errorf("seg %d: sport=%d want 12345", i, sport)
|
||||
}
|
||||
if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 53 {
|
||||
t.Errorf("seg %d: dport=%d want 53", i, dport)
|
||||
}
|
||||
// Software UDP GSO bumps the IPv4 ID per segment just like TSO
|
||||
// (inet_gso_segment's fixed-ID case is TCP-only).
|
||||
if id := binary.BigEndian.Uint16(seg[4:6]); id != 0x4242+uint16(i) {
|
||||
t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242+uint16(i))
|
||||
}
|
||||
|
||||
segPayLen := len(seg) - int(hdrLen)
|
||||
if udpLen := binary.BigEndian.Uint16(seg[24:26]); udpLen != uint16(8+segPayLen) {
|
||||
t.Errorf("seg %d: udp len=%d want %d", i, udpLen, 8+segPayLen)
|
||||
}
|
||||
|
||||
wantPay := make([]byte, segPayLen)
|
||||
for k := 0; k < segPayLen; k++ {
|
||||
wantPay[k] = byte((off + k) & 0xff)
|
||||
}
|
||||
if !bytes.Equal(seg[hdrLen:], wantPay) {
|
||||
t.Errorf("seg %d: payload mismatch", i)
|
||||
}
|
||||
off += segPayLen
|
||||
|
||||
if !verifyChecksum(seg[:20], 0) {
|
||||
t.Errorf("seg %d: bad IPv4 header checksum", i)
|
||||
}
|
||||
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_UDP, len(seg)-20)
|
||||
if !verifyChecksum(seg[20:], psum) {
|
||||
t.Errorf("seg %d: bad UDP checksum", i)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// buildUDPv4Single constructs a single IPv4/UDP datagram with the checksum field pre-loaded with the folded
|
||||
// pseudo-header sum, which is how a NEEDS_CSUM (CHECKSUM_PARTIAL) packet arrives from the tun.
|
||||
func buildUDPv4Single(payload []byte) (pkt []byte, hdr Hdr) {
|
||||
const ipLen, udpLen = 20, 8
|
||||
pkt = make([]byte, ipLen+udpLen+len(payload))
|
||||
|
||||
pkt[0] = 0x45
|
||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(len(pkt)))
|
||||
pkt[8] = 64
|
||||
pkt[9] = unix.IPPROTO_UDP
|
||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||
|
||||
binary.BigEndian.PutUint16(pkt[ipLen:ipLen+2], 12345)
|
||||
binary.BigEndian.PutUint16(pkt[ipLen+2:ipLen+4], 53)
|
||||
binary.BigEndian.PutUint16(pkt[ipLen+4:ipLen+6], uint16(udpLen+len(payload)))
|
||||
copy(pkt[ipLen+udpLen:], payload)
|
||||
|
||||
pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, udpLen+len(payload))
|
||||
binary.BigEndian.PutUint16(pkt[ipLen+udpChecksumOff:ipLen+udpChecksumOff+2], pseudo)
|
||||
|
||||
return pkt, NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, ipLen, udpChecksumOff)
|
||||
}
|
||||
|
||||
// TestFinishChecksumUDPZeroStoresAllOnes pins RFC 768: a UDP checksum that computes to zero goes on the wire as
|
||||
// 0xffff, because all-zero is the reserved "no checksum" encoding that IPv6 rejects outright.
|
||||
func TestFinishChecksumUDPZeroStoresAllOnes(t *testing.T) {
|
||||
var payload []byte
|
||||
for i := 0; i < 0x10000; i++ {
|
||||
p := []byte{byte(i >> 8), byte(i)}
|
||||
pkt, hdr := buildUDPv4Single(p)
|
||||
cs, co := int(hdr.CsumStart), int(hdr.CsumOffset)
|
||||
partial := binary.BigEndian.Uint16(pkt[cs+co : cs+co+2])
|
||||
pkt[cs+co], pkt[cs+co+1] = 0, 0
|
||||
if ^checksum.Checksum(pkt[cs:], partial) == 0 {
|
||||
payload = p
|
||||
break
|
||||
}
|
||||
}
|
||||
if payload == nil {
|
||||
t.Fatal("no 2-byte payload produced a zero checksum")
|
||||
}
|
||||
|
||||
pkt, hdr := buildUDPv4Single(payload)
|
||||
if err := FinishChecksum(pkt, hdr); err != nil {
|
||||
t.Fatalf("FinishChecksum: %v", err)
|
||||
}
|
||||
off := int(hdr.CsumStart) + int(hdr.CsumOffset)
|
||||
if got := binary.BigEndian.Uint16(pkt[off : off+2]); got != 0xffff {
|
||||
t.Fatalf("stored %#04x, want 0xffff for a UDP checksum that folds to zero", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFinishChecksumTCPZeroPreserved is the other half of the protocol split: zero is a legal TCP checksum and must
|
||||
// be stored as-is, so the UDP rewrite must not fire on a TCP csum_offset.
|
||||
func TestFinishChecksumTCPZeroPreserved(t *testing.T) {
|
||||
const cs, co = 20, tcpChecksumOff
|
||||
|
||||
// Seed the partial so the completed checksum lands on zero, the case the UDP path rewrites.
|
||||
seg := make([]byte, cs+co+2)
|
||||
for i := range seg[cs:] {
|
||||
seg[cs+i] = byte(i * 7)
|
||||
}
|
||||
var partial uint16
|
||||
for i := 0; i <= 0xffff; i++ {
|
||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], uint16(i))
|
||||
probe := append([]byte(nil), seg...)
|
||||
probe[cs+co], probe[cs+co+1] = 0, 0
|
||||
if ^checksum.Checksum(probe[cs:], uint16(i)) == 0 {
|
||||
partial = uint16(i)
|
||||
break
|
||||
}
|
||||
}
|
||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], partial)
|
||||
|
||||
hdr := NewHeader(unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, unix.VIRTIO_NET_HDR_GSO_NONE, 0, 0, cs, co)
|
||||
if err := FinishChecksum(seg, hdr); err != nil {
|
||||
t.Fatalf("FinishChecksum: %v", err)
|
||||
}
|
||||
if got := binary.BigEndian.Uint16(seg[cs+co : cs+co+2]); got != 0 {
|
||||
t.Fatalf("stored %#04x, want 0x0000 (zero is a legal TCP checksum)", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFinishChecksumUDPValidates confirms the ordinary path still produces a checksum a receiver accepts.
|
||||
func TestFinishChecksumUDPValidates(t *testing.T) {
|
||||
payload := []byte("the definitive tun offloads branch")
|
||||
pkt, hdr := buildUDPv4Single(payload)
|
||||
if err := FinishChecksum(pkt, hdr); err != nil {
|
||||
t.Fatalf("FinishChecksum: %v", err)
|
||||
}
|
||||
pseudo := pseudoHeaderIPv4(pkt[12:16], pkt[16:20], unix.IPPROTO_UDP, 8+len(payload))
|
||||
if !verifyChecksum(pkt[hdr.CsumStart:], pseudo) {
|
||||
t.Fatal("completed UDP checksum does not validate")
|
||||
}
|
||||
}
|
||||
|
||||
// TestCheckValidMasksGSOECN: GSO_ECN is a qualifier bit the kernel ORs
|
||||
// into gso_type for TSO superpackets with CWR set. CheckValid must
|
||||
// validate an ECN-qualified type as its base type — previously TCPV4|ECN
|
||||
// fell into the default case and skipped the IP-version agreement check.
|
||||
// The qualifier is TCP-only, so it must be rejected on UDP_L4.
|
||||
func TestCheckValidMasksGSOECN(t *testing.T) {
|
||||
v4pkt, _, _ := buildTCPv4Super(100)
|
||||
v6pkt := make([]byte, len(v4pkt))
|
||||
copy(v6pkt, v4pkt)
|
||||
v6pkt[0] = 0x60 // claim IPv6
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
pkt []byte
|
||||
gsoType uint8
|
||||
wantErr bool
|
||||
}{
|
||||
{"tcpv4-ecn-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, false},
|
||||
{"tcpv4-ecn-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
|
||||
{"tcpv6-ecn-v4-mismatch", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV6 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
|
||||
{"udp-l4-ecn-rejected", v4pkt, unix.VIRTIO_NET_HDR_GSO_UDP_L4 | unix.VIRTIO_NET_HDR_GSO_ECN, true},
|
||||
{"tcpv4-plain-v4", v4pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, false},
|
||||
{"tcpv4-plain-v6-mismatch", v6pkt, unix.VIRTIO_NET_HDR_GSO_TCPV4, true},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
err := CheckValid(tc.pkt, NewHeader(0, tc.gsoType, 0, 100, 0, 0))
|
||||
if tc.wantErr && err == nil {
|
||||
t.Errorf("CheckValid(gsoType=%#x) = nil, want error", tc.gsoType)
|
||||
}
|
||||
if !tc.wantErr && err != nil {
|
||||
t.Errorf("CheckValid(gsoType=%#x) = %v, want nil", tc.gsoType, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestCheckValidRejectsZeroGSOSize: a GSO-typed header with gso_size=0 must be
|
||||
// rejected. It would produce a Packet whose GSOInfo.IsSuperpacket() is false,
|
||||
// dodging both segmentation and FinishChecksum on its way downstream.
|
||||
func TestCheckValidRejectsZeroGSOSize(t *testing.T) {
|
||||
v4pkt, _, _ := buildTCPv4Super(100)
|
||||
if err := CheckValid(v4pkt, NewHeader(0, unix.VIRTIO_NET_HDR_GSO_TCPV4, 0, 0, 0, 0)); err == nil {
|
||||
t.Fatal("CheckValid accepted a GSO-typed header with gso_size=0")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFoldComplementMatchesReference checks the segmenter's fold-and-invert
|
||||
// against an independent RFC 1071 reference fold, hitting the carry edge
|
||||
// cases (values whose first fold produces another carry).
|
||||
func TestFoldComplementMatchesReference(t *testing.T) {
|
||||
refFold := func(s uint64) uint16 {
|
||||
for s>>16 != 0 {
|
||||
s = s&0xffff + s>>16
|
||||
}
|
||||
return uint16(s)
|
||||
}
|
||||
cases := []uint32{
|
||||
0, 1, 0xffff,
|
||||
0x10000, // single carry
|
||||
0x1fffe, // 0xffff + 0xffff: carry produces another 0xffff
|
||||
0xffff0000, // high half only
|
||||
0xfffeffff, // first fold yields another carry
|
||||
0xffffffff, // worst case
|
||||
}
|
||||
for _, c := range cases {
|
||||
if got, want := foldComplement(c), ^refFold(uint64(c)); got != want {
|
||||
t.Errorf("foldComplement(%#x) = %#x, want %#x", c, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// referenceBaseIPv4HdrSum and referenceBaseTCPHdrSum are the straightforward
|
||||
// implementations that baseIPv4HdrSum/baseTCPHdrSum replaced: copy the header
|
||||
// into scratch, zero the fields the segment loop rewrites, sum. The production
|
||||
// versions instead sum in place and subtract those fields via one's-complement
|
||||
// arithmetic, which is faster but far less obvious — particularly for the TCP
|
||||
// flags byte, which is only half of a 16-bit word. These references exist so
|
||||
// that trade is checked rather than asserted.
|
||||
func referenceBaseIPv4HdrSum(pkt []byte, ihl int) uint32 {
|
||||
var ipTmp [ipv4HeaderMaxLen]byte
|
||||
copy(ipTmp[:ihl], pkt[:ihl])
|
||||
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
|
||||
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
|
||||
ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0
|
||||
return uint32(checksum.Checksum(ipTmp[:ihl], 0))
|
||||
}
|
||||
|
||||
func referenceBaseTCPHdrSum(pkt []byte, csumStart, headerLen int) uint32 {
|
||||
tcpLen := headerLen - csumStart
|
||||
var tmp [tcpHeaderMaxLen]byte
|
||||
copy(tmp[:tcpLen], pkt[csumStart:headerLen])
|
||||
tmp[tcpSeqOff], tmp[tcpSeqOff+1], tmp[tcpSeqOff+2], tmp[tcpSeqOff+3] = 0, 0, 0, 0
|
||||
tmp[tcpFlagsOff] = 0
|
||||
tmp[tcpChecksumOff], tmp[tcpChecksumOff+1] = 0, 0
|
||||
return uint32(checksum.Checksum(tmp[:tcpLen], 0))
|
||||
}
|
||||
|
||||
// randSeed is a tiny deterministic PRNG so this test needs no imports beyond
|
||||
// what the file already has and reproduces identically on every run.
|
||||
func randByte(state *uint32) byte {
|
||||
*state = *state*1664525 + 1013904223
|
||||
return byte(*state >> 24)
|
||||
}
|
||||
|
||||
func TestBaseSumsMatchZeroingReference(t *testing.T) {
|
||||
state := uint32(12345)
|
||||
|
||||
t.Run("ipv4", func(t *testing.T) {
|
||||
for ihl := ipv4HeaderMinLen; ihl <= ipv4HeaderMaxLen; ihl += 4 {
|
||||
for iter := 0; iter < 5000; iter++ {
|
||||
pkt := make([]byte, ihl)
|
||||
for i := range pkt {
|
||||
pkt[i] = randByte(&state)
|
||||
}
|
||||
pkt[0] = byte(0x40 | (ihl / 4))
|
||||
|
||||
want := referenceBaseIPv4HdrSum(pkt, ihl)
|
||||
got, err := baseIPv4HdrSum(pkt, ihl)
|
||||
if err != nil {
|
||||
t.Fatalf("ihl=%d: %v", ihl, err)
|
||||
}
|
||||
// Compare the value that reaches the wire: the raw partial
|
||||
// sums may legally differ by one's-complement -0 vs +0.
|
||||
for _, tl := range []uint32{20, 1500, 65535} {
|
||||
for _, id := range []uint32{0, 0x4242, 0xffff} {
|
||||
if a, b := foldComplement(want+tl+id), foldComplement(got+tl+id); a != b {
|
||||
t.Fatalf("ihl=%d tl=%d id=%d: %#04x != %#04x", ihl, tl, id, a, b)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("tcp", func(t *testing.T) {
|
||||
const csumStart = 20
|
||||
for dataOff := 5; dataOff <= 15; dataOff++ {
|
||||
tcpLen := dataOff * 4
|
||||
headerLen := csumStart + tcpLen
|
||||
for iter := 0; iter < 5000; iter++ {
|
||||
pkt := make([]byte, headerLen+64)
|
||||
for i := range pkt {
|
||||
pkt[i] = randByte(&state)
|
||||
}
|
||||
pkt[0] = 0x45
|
||||
pkt[csumStart+tcpDataOffOff] = byte(dataOff << 4)
|
||||
|
||||
want := referenceBaseTCPHdrSum(pkt, csumStart, headerLen)
|
||||
got := baseTCPHdrSum(pkt, csumStart, headerLen)
|
||||
for _, seq := range []uint32{0, 1, 0x4242_4242, 0xffff_ffff} {
|
||||
for _, fl := range []uint32{0x00, 0x10, 0x18, 0x19, 0xff} {
|
||||
for _, l4 := range []uint32{20, 1460, 65535} {
|
||||
a := foldComplement(want + seq + fl + l4)
|
||||
b := foldComplement(got + seq + fl + l4)
|
||||
if a != b {
|
||||
t.Fatalf("dataOff=%d seq=%#x fl=%#x l4=%d: %#04x != %#04x",
|
||||
dataOff, seq, fl, l4, a, b)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user