mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-16 08:46:58 +02:00
udp: test the partial-sendmmsg rewind
The rewind (resume after the kernel accepts fewer entries than submitted) was the hairiest untested logic in the write path; a bug there silently duplicates or loses packets under backpressure. Give batchWriter an injectable sendFn and drive WriteBatch through scripted partial-acceptance sequences over a mixed GSO-run/plain batch, decoding what "reached the wire" straight from the prepared iovecs rather than the entryEnd bookkeeping under test. Also pins the zero-progress abort and the EIO runtime GSO-disable replay. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -4,6 +4,7 @@ package udp
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
|
"fmt"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -55,10 +55,16 @@ type batchWriter struct {
|
|||||||
// entryEnd[e] is the bufs index *after* the last packet packed into
|
// entryEnd[e] is the bufs index *after* the last packet packed into
|
||||||
// mmsghdr entry e. Used to rewind `i` on partial sendmmsg success.
|
// mmsghdr entry e. Used to rewind `i` on partial sendmmsg success.
|
||||||
entryEnd []int
|
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 {
|
func newBatchWriter(fd int, isV4 bool, l *slog.Logger) *batchWriter {
|
||||||
w := &batchWriter{fd: fd, isV4: isV4, l: l}
|
w := &batchWriter{fd: fd, isV4: isV4, l: l}
|
||||||
|
w.sendFn = w.sendmmsg
|
||||||
w.prepareWriteMessages(MaxWriteBatch)
|
w.prepareWriteMessages(MaxWriteBatch)
|
||||||
w.prepareGSO()
|
w.prepareGSO()
|
||||||
return w
|
return w
|
||||||
@@ -294,7 +300,7 @@ sendChunks:
|
|||||||
return written, fmt.Errorf("sendmmsg: no progress")
|
return written, fmt.Errorf("sendmmsg: no progress")
|
||||||
}
|
}
|
||||||
|
|
||||||
sent, serr := w.sendmmsg(entry)
|
sent, serr := w.sendFn(entry)
|
||||||
if serr != nil && sent <= 0 {
|
if serr != nil && sent <= 0 {
|
||||||
// sent<=0 means message 0 itself failed. If that entry was a GSO
|
// 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"
|
// superpacket and the errno is the kernel's "device can't do this"
|
||||||
|
|||||||
Reference in New Issue
Block a user