From 66cb98e13a53fc57e0c773ed949ae65cd3d55a34 Mon Sep 17 00:00:00 2001 From: JackDoan Date: Mon, 27 Jul 2026 15:59:49 -0500 Subject: [PATCH] udp: test the partial-sendmmsg rewind The rewind (resume after the kernel accepts fewer entries than submitted) was the hairiest untested logic in the write path; a bug there silently duplicates or loses packets under backpressure. Give batchWriter an injectable sendFn and drive WriteBatch through scripted partial-acceptance sequences over a mixed GSO-run/plain batch, decoding what "reached the wire" straight from the prepared iovecs rather than the entryEnd bookkeeping under test. Also pins the zero-progress abort and the EIO runtime GSO-disable replay. Co-Authored-By: Claude Fable 5 --- udp/udp_linux_fixes_test.go | 153 ++++++++++++++++++++++++++++++++++++ udp/udp_linux_writebatch.go | 8 +- 2 files changed, 160 insertions(+), 1 deletion(-) diff --git a/udp/udp_linux_fixes_test.go b/udp/udp_linux_fixes_test.go index b6e9ffa1..ce25ba70 100644 --- a/udp/udp_linux_fixes_test.go +++ b/udp/udp_linux_fixes_test.go @@ -4,6 +4,7 @@ package udp import ( "encoding/binary" + "fmt" "log/slog" "net" "net/netip" @@ -397,3 +398,155 @@ func TestDeliverSegments(t *testing.T) { }) } } + +// newRewindTestWriter builds a batchWriter with no socket: GSO planning on, +// sendFn left for the test to script. fd is invalid on purpose -- any path +// that actually hits the kernel (sendto fallback) fails loudly. +func newRewindTestWriter() *batchWriter { + w := &batchWriter{fd: -1, isV4: true, l: testLogger()} + w.prepareWriteMessages(MaxWriteBatch) + w.gsoSupported = true + w.maxGSOSegments = 63 + return w +} + +// capturePrepared decodes the first n prepared mmsghdr entries straight +// from their iovecs -- ground truth, deliberately not the entryEnd +// bookkeeping the rewind logic itself relies on. Returns one []byte per +// packed packet, in entry order. +func capturePrepared(w *batchWriter, n int) [][]byte { + var out [][]byte + for e := 0; e < n; e++ { + hdr := &w.msgs[e].Hdr + iovs := unsafe.Slice(hdr.Iov, int(hdr.Iovlen)) + for _, iov := range iovs { + b := make([]byte, int(iov.Len)) + if iov.Len > 0 { + copy(b, unsafe.Slice(iov.Base, int(iov.Len))) + } + out = append(out, b) + } + } + return out +} + +// TestWriteBatchPartialSendRewind drives WriteBatch through scripted +// partial sendmmsg results and asserts the rewind resumes exactly where +// the kernel stopped: every packet on the wire exactly once, in order, +// no duplicate, no loss. This is the hairiest logic in the write path +// and a rewind bug means silent packet duplication or loss under EAGAIN- +// style backpressure. +func TestWriteBatchPartialSendRewind(t *testing.T) { + dstA := netip.MustParseAddrPort("127.0.0.1:4242") + dstB := netip.MustParseAddrPort("127.0.0.2:4242") + + mkBuf := func(tag byte, n int) []byte { + b := make([]byte, n) + for i := range b { + b[i] = tag + } + b[0] = tag // tag identifies the packet uniquely below + return b + } + + // Mixed shape: a 3-packet GSO run to A, a lone short packet to A (run + // tail), then two to B. The planner packs this as multiple entries with + // multi-iovec runs, which is what makes the rewind arithmetic hairy. + bufs := [][]byte{ + mkBuf(1, 1200), mkBuf(2, 1200), mkBuf(3, 1200), // run to A + mkBuf(4, 600), // short tail to A + mkBuf(5, 900), mkBuf(6, 900), // run to B + } + addrs := []netip.AddrPort{dstA, dstA, dstA, dstA, dstB, dstB} + + scripts := [][]int{ + {99}, // accept everything first call + {1, 99}, // one entry per call, then the rest + {1, 1, 1, 99}, // strictly one entry per call + {2, 99}, // two entries, then the rest + } + for si, script := range scripts { + t.Run(fmt.Sprintf("script_%d", si), func(t *testing.T) { + w := newRewindTestWriter() + var wire [][]byte + call := 0 + w.sendFn = func(n int) (int, error) { + accept := n + if call < len(script) && script[call] < n { + accept = script[call] + } + call++ + wire = append(wire, capturePrepared(w, accept)...) + return accept, nil + } + + written, err := w.WriteBatch(bufs, addrs, nil) + if err != nil { + t.Fatalf("WriteBatch: %v", err) + } + if written != len(bufs) { + t.Errorf("written = %d, want %d", written, len(bufs)) + } + if len(wire) != len(bufs) { + t.Fatalf("wire got %d packets, want %d (dup or loss in rewind)", len(wire), len(bufs)) + } + for i, b := range wire { + if len(b) != len(bufs[i]) || b[0] != bufs[i][0] { + t.Errorf("wire[%d] = tag %d len %d, want tag %d len %d (reorder/dup)", + i, b[0], len(b), bufs[i][0], len(bufs[i])) + } + } + }) + } +} + +// TestWriteBatchZeroProgress: sent == 0 with no error must abort with an +// error rather than spin forever replaying the same chunk. +func TestWriteBatchZeroProgress(t *testing.T) { + w := newRewindTestWriter() + w.sendFn = func(n int) (int, error) { return 0, nil } + bufs := [][]byte{make([]byte, 100)} + addrs := []netip.AddrPort{netip.MustParseAddrPort("127.0.0.1:4242")} + if _, err := w.WriteBatch(bufs, addrs, nil); err == nil { + t.Fatal("WriteBatch = nil error on zero progress, want error") + } +} + +// TestWriteBatchEIODisablesGSOAndReplays pins the runtime GSO give-up: a +// sendmmsg rejected with EIO on a GSO superpacket entry must clear +// gsoSupported and replay the same packets as per-packet entries through +// sendmmsg (keeping batching), not fall back to per-packet sendto. +func TestWriteBatchEIODisablesGSOAndReplays(t *testing.T) { + dst := netip.MustParseAddrPort("127.0.0.1:4242") + bufs := [][]byte{make([]byte, 1200), make([]byte, 1200), make([]byte, 1200)} + addrs := []netip.AddrPort{dst, dst, dst} + + w := newRewindTestWriter() + var entryCounts []int + call := 0 + w.sendFn = func(n int) (int, error) { + entryCounts = append(entryCounts, n) + call++ + if call == 1 { + return -1, &net.OpError{Op: "sendmmsg", Err: unix.EIO} + } + return n, nil + } + + written, err := w.WriteBatch(bufs, addrs, nil) + if err != nil { + t.Fatalf("WriteBatch: %v", err) + } + if w.gsoSupported { + t.Error("gsoSupported still true after EIO on a GSO entry") + } + if written != len(bufs) { + t.Errorf("written = %d, want %d", written, len(bufs)) + } + // First call: one GSO entry carrying the whole run. Replay: one entry + // per packet, still via sendmmsg. + want := []int{1, 3} + if len(entryCounts) != len(want) || entryCounts[0] != want[0] || entryCounts[1] != want[1] { + t.Errorf("sendmmsg entry counts = %v, want %v", entryCounts, want) + } +} diff --git a/udp/udp_linux_writebatch.go b/udp/udp_linux_writebatch.go index c807ffa5..480ce413 100644 --- a/udp/udp_linux_writebatch.go +++ b/udp/udp_linux_writebatch.go @@ -55,10 +55,16 @@ type batchWriter struct { // entryEnd[e] is the bufs index *after* the last packet packed into // mmsghdr entry e. Used to rewind `i` on partial sendmmsg success. entryEnd []int + + // sendFn issues the sendmmsg for the first n prepared entries. Points + // at the real syscall in production; tests inject partial-success and + // error scripts to exercise the rewind logic without a socket. + sendFn func(n int) (int, error) } func newBatchWriter(fd int, isV4 bool, l *slog.Logger) *batchWriter { w := &batchWriter{fd: fd, isV4: isV4, l: l} + w.sendFn = w.sendmmsg w.prepareWriteMessages(MaxWriteBatch) w.prepareGSO() return w @@ -294,7 +300,7 @@ sendChunks: return written, fmt.Errorf("sendmmsg: no progress") } - sent, serr := w.sendmmsg(entry) + sent, serr := w.sendFn(entry) if serr != nil && sent <= 0 { // sent<=0 means message 0 itself failed. If that entry was a GSO // superpacket and the errno is the kernel's "device can't do this"