mirror of
https://github.com/slackhq/nebula.git
synced 2026-10-08 10:27:55 +02:00
the definitive tun offloads branch (#1704)
This commit is contained in:
+30
-2
@@ -8,16 +8,41 @@ import (
|
||||
|
||||
const MTU = 9001
|
||||
|
||||
// MaxWriteBatch is the largest batch any Conn.WriteBatch implementation is
|
||||
// required to accept. Callers SHOULD NOT pass more than this per call; Linux
|
||||
// backends preallocate sendmmsg scratch sized to this value, so exceeding it
|
||||
// only costs additional sendmmsg chunks within a single WriteBatch call.
|
||||
const MaxWriteBatch = 128
|
||||
|
||||
type EncReader func(
|
||||
addr netip.AddrPort,
|
||||
payload []byte,
|
||||
)
|
||||
|
||||
type Settings struct {
|
||||
Listen netip.AddrPort
|
||||
Multi bool
|
||||
Batch int
|
||||
Offloads bool
|
||||
}
|
||||
|
||||
type Conn interface {
|
||||
Rebind() error
|
||||
LocalAddr() (netip.AddrPort, error)
|
||||
ListenOut(r EncReader) error
|
||||
// ListenOut invokes r for each received packet.
|
||||
// On batch-capable backends (recvmmsg), flush is called after each batch is fully delivered.
|
||||
// Callers use it to flush per-batch accumulators such as TUN write coalescers.
|
||||
// Single-packet backends call flush after each packet. flush must not be nil.
|
||||
ListenOut(r EncReader, flush func()) error
|
||||
WriteTo(b []byte, addr netip.AddrPort) error
|
||||
// WriteBatch sends a contiguous batch of packets, each with its own
|
||||
// destination. bufs and addrs must have the same length. Linux uses
|
||||
// sendmmsg(2) for a single syscall.
|
||||
//
|
||||
// Returns the number of packets successfully written. A destination the kernel rejects costs only
|
||||
// its own packet, so a short count means some peers were undeliverable, not that the batch failed.
|
||||
// Not safe for concurrent use on the same Conn.
|
||||
WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error)
|
||||
ReloadConfig(c *config.C)
|
||||
SupportsMultipleReaders() bool
|
||||
Close() error
|
||||
@@ -31,7 +56,7 @@ func (NoopConn) Rebind() error {
|
||||
func (NoopConn) LocalAddr() (netip.AddrPort, error) {
|
||||
return netip.AddrPort{}, nil
|
||||
}
|
||||
func (NoopConn) ListenOut(_ EncReader) error {
|
||||
func (NoopConn) ListenOut(_ EncReader, _ func()) error {
|
||||
return nil
|
||||
}
|
||||
func (NoopConn) SupportsMultipleReaders() bool {
|
||||
@@ -40,6 +65,9 @@ func (NoopConn) SupportsMultipleReaders() bool {
|
||||
func (NoopConn) WriteTo(_ []byte, _ netip.AddrPort) error {
|
||||
return nil
|
||||
}
|
||||
func (NoopConn) WriteBatch(bufs [][]byte, _ []netip.AddrPort) (int, error) {
|
||||
return len(bufs), nil
|
||||
}
|
||||
func (NoopConn) ReloadConfig(_ *config.C) {
|
||||
return
|
||||
}
|
||||
|
||||
+2
-3
@@ -7,14 +7,13 @@ import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"syscall"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) {
|
||||
return NewGenericListener(l, ip, port, multi, batch)
|
||||
func NewListener(l *slog.Logger, s Settings) (Conn, error) {
|
||||
return NewGenericListener(l, s)
|
||||
}
|
||||
|
||||
func NewListenConfig(multi bool) net.ListenConfig {
|
||||
|
||||
+2
-3
@@ -10,14 +10,13 @@ import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"syscall"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) {
|
||||
return NewGenericListener(l, ip, port, multi, batch)
|
||||
func NewListener(l *slog.Logger, s Settings) (Conn, error) {
|
||||
return NewGenericListener(l, s)
|
||||
}
|
||||
|
||||
func NewListenConfig(multi bool) net.ListenConfig {
|
||||
|
||||
+22
-5
@@ -27,9 +27,9 @@ type StdConn struct {
|
||||
|
||||
var _ Conn = &StdConn{}
|
||||
|
||||
func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) {
|
||||
lc := NewListenConfig(multi)
|
||||
pc, err := lc.ListenPacket(context.TODO(), "udp", net.JoinHostPort(ip.String(), fmt.Sprintf("%v", port)))
|
||||
func NewListener(l *slog.Logger, s Settings) (Conn, error) {
|
||||
lc := NewListenConfig(s.Multi)
|
||||
pc, err := lc.ListenPacket(context.TODO(), "udp", s.Listen.String())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -140,6 +140,22 @@ func (u *StdConn) WriteTo(b []byte, ap netip.AddrPort) error {
|
||||
}
|
||||
}
|
||||
|
||||
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
|
||||
// An un-sendable destination costs its own packet, never the ones behind it in the batch.
|
||||
// TODO: WriteTo maps EWOULDBLOCK to an error, so a full send buffer
|
||||
// silently drops the rest of a burst (linux blocks instead). Poll for
|
||||
// writability on EAGAIN before giving up on the remainder.
|
||||
written := 0
|
||||
for i, b := range bufs {
|
||||
if err := u.WriteTo(b, addrs[i]); err == nil {
|
||||
written++
|
||||
} else {
|
||||
u.l.Debug("failed to write packet in batch", "udpAddr", addrs[i], "error", err)
|
||||
}
|
||||
}
|
||||
return written, nil
|
||||
}
|
||||
|
||||
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
||||
a := u.UDPConn.LocalAddr()
|
||||
|
||||
@@ -165,7 +181,7 @@ func NewUDPStatsEmitter(udpConns []Conn) func() {
|
||||
return func() {}
|
||||
}
|
||||
|
||||
func (u *StdConn) ListenOut(r EncReader) error {
|
||||
func (u *StdConn) ListenOut(r EncReader, flush func()) error {
|
||||
buffer := make([]byte, MTU)
|
||||
|
||||
for {
|
||||
@@ -179,7 +195,8 @@ func (u *StdConn) ListenOut(r EncReader) error {
|
||||
continue
|
||||
}
|
||||
|
||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n])
|
||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n:n])
|
||||
flush()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+19
-5
@@ -27,9 +27,9 @@ type GenericConn struct {
|
||||
|
||||
var _ Conn = &GenericConn{}
|
||||
|
||||
func NewGenericListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) {
|
||||
lc := NewListenConfig(multi)
|
||||
pc, err := lc.ListenPacket(context.TODO(), "udp", net.JoinHostPort(ip.String(), fmt.Sprintf("%v", port)))
|
||||
func NewGenericListener(l *slog.Logger, s Settings) (Conn, error) {
|
||||
lc := NewListenConfig(s.Multi)
|
||||
pc, err := lc.ListenPacket(context.TODO(), "udp", s.Listen.String())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -44,6 +44,19 @@ func (u *GenericConn) WriteTo(b []byte, addr netip.AddrPort) error {
|
||||
return err
|
||||
}
|
||||
|
||||
func (u *GenericConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
|
||||
// An un-sendable destination costs its own packet, never the ones behind it in the batch.
|
||||
written := 0
|
||||
for i, b := range bufs {
|
||||
if _, err := u.UDPConn.WriteToUDPAddrPort(b, addrs[i]); err == nil {
|
||||
written++
|
||||
} else {
|
||||
u.l.Debug("failed to write packet in batch", "udpAddr", addrs[i], "error", err)
|
||||
}
|
||||
}
|
||||
return written, nil
|
||||
}
|
||||
|
||||
func (u *GenericConn) LocalAddr() (netip.AddrPort, error) {
|
||||
a := u.UDPConn.LocalAddr()
|
||||
|
||||
@@ -73,7 +86,7 @@ type rawMessage struct {
|
||||
Len uint32
|
||||
}
|
||||
|
||||
func (u *GenericConn) ListenOut(r EncReader) error {
|
||||
func (u *GenericConn) ListenOut(r EncReader, flush func()) error {
|
||||
buffer := make([]byte, MTU)
|
||||
|
||||
var lastRecvErr time.Time
|
||||
@@ -93,7 +106,8 @@ func (u *GenericConn) ListenOut(r EncReader) error {
|
||||
continue
|
||||
}
|
||||
|
||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n])
|
||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n:n])
|
||||
flush()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+234
-89
@@ -1,5 +1,4 @@
|
||||
//go:build !android && !e2e_testing
|
||||
// +build !android,!e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
@@ -25,11 +24,18 @@ type StdConn struct {
|
||||
isV4 bool
|
||||
l *slog.Logger
|
||||
batch int
|
||||
|
||||
// bw owns the sendmmsg/UDP-GSO transmit path: the per-queue write
|
||||
// scratch and the GSO capability state probed at socket creation. See
|
||||
// udp_linux_writebatch.go.
|
||||
bw *batchWriter
|
||||
|
||||
groSupported bool
|
||||
}
|
||||
|
||||
func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) {
|
||||
func NewListener(l *slog.Logger, s Settings) (Conn, error) {
|
||||
af := unix.AF_INET6
|
||||
if ip.Is4() {
|
||||
if s.Listen.Addr().Is4() {
|
||||
af = unix.AF_INET
|
||||
}
|
||||
syscall.ForkLock.RLock()
|
||||
@@ -42,7 +48,7 @@ func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int)
|
||||
return nil, fmt.Errorf("unable to open socket: %w", err)
|
||||
}
|
||||
|
||||
if multi {
|
||||
if s.Multi {
|
||||
if err = unix.SetsockoptInt(fd, unix.SOL_SOCKET, unix.SO_REUSEPORT, 1); err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return nil, fmt.Errorf("unable to set SO_REUSEPORT: %w", err)
|
||||
@@ -50,13 +56,14 @@ func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int)
|
||||
}
|
||||
|
||||
var sa unix.Sockaddr
|
||||
if ip.Is4() {
|
||||
port := int(s.Listen.Port())
|
||||
if s.Listen.Addr().Is4() {
|
||||
sa4 := &unix.SockaddrInet4{Port: port}
|
||||
sa4.Addr = ip.As4()
|
||||
sa4.Addr = s.Listen.Addr().As4()
|
||||
sa = sa4
|
||||
} else {
|
||||
sa6 := &unix.SockaddrInet6{Port: port}
|
||||
sa6.Addr = ip.As16()
|
||||
sa6.Addr = s.Listen.Addr().As16()
|
||||
sa = sa6
|
||||
}
|
||||
if err = unix.Bind(fd, sa); err != nil {
|
||||
@@ -64,7 +71,60 @@ func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int)
|
||||
return nil, fmt.Errorf("unable to bind to socket: %w", err)
|
||||
}
|
||||
|
||||
return &StdConn{sysFd: fd, isV4: ip.Is4(), l: l, batch: batch}, nil
|
||||
out := &StdConn{sysFd: fd, isV4: s.Listen.Addr().Is4(), l: l, batch: s.Batch}
|
||||
|
||||
out.bw = newBatchWriter(fd, out.isV4, l, s.Offloads)
|
||||
|
||||
// GRO coalesces same-flow datagrams into superpackets that must be split back apart via the delivered gso_size cmsg
|
||||
// batch == 1 means the caller wants plain single-datagram reads with MTU-sized buffers, so leave it off.
|
||||
if s.Batch > 1 && s.Offloads {
|
||||
out.prepareGRO()
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// udpGROBufferSize sizes the per-entry recvmmsg buffer when UDP_GRO is on.
|
||||
// The kernel stitches a run of same-flow datagrams into a single skb whose
|
||||
// length is bounded by sk_gso_max_size (65535)
|
||||
const udpGROBufferSize = 65535
|
||||
|
||||
// udpGROCmsgPayload is the size of the UDP_GRO cmsg data delivered by the
|
||||
// kernel: a single int (gso_size in bytes). See udp_cmsg_recv() in net/ipv4/udp.c.
|
||||
const udpGROCmsgPayload = 4
|
||||
|
||||
// prepareGRO turns on UDP_GRO so the kernel coalesces consecutive same-flow
|
||||
// datagrams into one recvmmsg entry, with a cmsg carrying the gso_size used
|
||||
// to split them back apart on the application side.
|
||||
func (u *StdConn) prepareGRO() {
|
||||
err := unix.SetsockoptInt(u.sysFd, unix.IPPROTO_UDP, unix.UDP_GRO, 1)
|
||||
if err != nil {
|
||||
u.l.Info("udp: GRO disabled", "reason", "kernel rejected probe", "error", err)
|
||||
recordCapability("udp.gro.enabled", false)
|
||||
return
|
||||
}
|
||||
u.groSupported = true
|
||||
u.l.Info("udp: GRO enabled")
|
||||
recordCapability("udp.gro.enabled", true)
|
||||
}
|
||||
|
||||
// recordCapability registers (or updates) a boolean gauge for one of the
|
||||
// kernel-feature probes. Gauges go to 1 when the feature is enabled, 0 when
|
||||
// it is not — dashboards can show degraded state on partially-supported
|
||||
// kernels at a glance. Calling repeatedly with the same name updates the
|
||||
// existing gauge rather than registering a duplicate.
|
||||
//
|
||||
// Caveat: the gauge is process-global while the capability state it reports
|
||||
// is per-socket. With multiple listen routines the last probe wins, and a
|
||||
// runtime downgrade on one socket (e.g. the GSO EIO disable) flips the gauge
|
||||
// for all of them. Treat it as "at least one socket looks like this."
|
||||
func recordCapability(name string, enabled bool) {
|
||||
g := metrics.GetOrRegisterGauge(name, nil)
|
||||
if enabled {
|
||||
g.Update(1)
|
||||
} else {
|
||||
g.Update(0)
|
||||
}
|
||||
}
|
||||
|
||||
func (u *StdConn) SupportsMultipleReaders() bool {
|
||||
@@ -114,7 +174,7 @@ func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
||||
}
|
||||
}
|
||||
|
||||
// recvmmsg does one blocking recvmmsg (MSG_WAITFORONE), reading up to len(msgs) datagrams
|
||||
// recvmmsg does one blocking recvmmsg (MSG_WAITFORONE), reading up to len(msgs) datagrams.
|
||||
func (u *StdConn) recvmmsg(msgs []rawMessage) (int, error) {
|
||||
r, _, errno := unix.Syscall6(
|
||||
unix.SYS_RECVMMSG,
|
||||
@@ -138,40 +198,70 @@ func (u *StdConn) recvmmsg(msgs []rawMessage) (int, error) {
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// recvmsg does one blocking recvmsg into msgs[0]
|
||||
func (u *StdConn) recvmsg(msgs []rawMessage) (int, error) {
|
||||
r, _, errno := unix.Syscall6(
|
||||
unix.SYS_RECVMSG,
|
||||
uintptr(u.sysFd),
|
||||
uintptr(unsafe.Pointer(&msgs[0].Hdr)),
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
)
|
||||
if errno != 0 {
|
||||
if u.closed.Load() {
|
||||
return 0, net.ErrClosed
|
||||
// prepareRawMessages allocates the recvmmsg scratch:
|
||||
// n rawMessages, each wired to its own bufSize receive buffer, sockaddr name slot
|
||||
// and, when cmsgSpace > 0, a slice of one contiguous ancillary-data slab.
|
||||
// All iovecs share a single slab kept alive by the msghdrs that point into it.
|
||||
func prepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
|
||||
msgs := make([]rawMessage, n)
|
||||
buffers := make([][]byte, n)
|
||||
names := make([][]byte, n)
|
||||
iovs := make([]iovec, n)
|
||||
|
||||
var cmsgs []byte
|
||||
if cmsgSpace > 0 {
|
||||
cmsgs = make([]byte, n*cmsgSpace)
|
||||
}
|
||||
|
||||
for i := range msgs {
|
||||
buffers[i] = make([]byte, bufSize)
|
||||
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||
|
||||
iovs[i].Base = &buffers[i][0]
|
||||
setIovLen(&iovs[i], bufSize)
|
||||
msgs[i].Hdr.Iov = &iovs[i]
|
||||
setMsgIovlen(&msgs[i].Hdr, 1)
|
||||
|
||||
msgs[i].Hdr.Name = &names[i][0]
|
||||
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
||||
|
||||
if cmsgSpace > 0 {
|
||||
msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace]
|
||||
setMsgControllen(&msgs[i].Hdr, cmsgSpace)
|
||||
}
|
||||
return 0, &net.OpError{Op: "recvmsg", Err: errno}
|
||||
}
|
||||
if r == 0 && u.closed.Load() {
|
||||
return 0, net.ErrClosed
|
||||
}
|
||||
msgs[0].Len = uint32(r)
|
||||
return 1, nil
|
||||
|
||||
return msgs, buffers, names, cmsgs
|
||||
}
|
||||
|
||||
func (u *StdConn) ListenOut(r EncReader) error {
|
||||
func getFrom(names [][]byte, i int, isV4 bool) netip.AddrPort {
|
||||
var ip netip.Addr
|
||||
msgs, buffers, names := u.PrepareRawMessages(u.batch)
|
||||
read := u.recvmmsg
|
||||
if u.batch == 1 {
|
||||
read = u.recvmsg
|
||||
// Its ok to skip the ok check here, the slicing is the only error that can occur and it will panic
|
||||
if isV4 {
|
||||
ip, _ = netip.AddrFromSlice(names[i][4:8])
|
||||
} else {
|
||||
ip, _ = netip.AddrFromSlice(names[i][8:24])
|
||||
}
|
||||
return netip.AddrPortFrom(ip.Unmap(), binary.BigEndian.Uint16(names[i][2:4]))
|
||||
}
|
||||
|
||||
func (u *StdConn) ListenOut(r EncReader, flush func()) error {
|
||||
bufSize := MTU
|
||||
cmsgSpace := 0
|
||||
if u.groSupported {
|
||||
bufSize = udpGROBufferSize
|
||||
cmsgSpace = unix.CmsgSpace(udpGROCmsgPayload)
|
||||
}
|
||||
msgs, buffers, names, _ := prepareRawMessages(u.batch, bufSize, cmsgSpace)
|
||||
|
||||
for {
|
||||
n, err := read(msgs)
|
||||
if cmsgSpace > 0 {
|
||||
for i := range msgs {
|
||||
setMsgControllen(&msgs[i].Hdr, cmsgSpace)
|
||||
}
|
||||
}
|
||||
|
||||
n, err := u.recvmmsg(msgs)
|
||||
if err != nil {
|
||||
if errors.Is(err, unix.EINTR) {
|
||||
continue // interrupted by a signal, retry the read
|
||||
@@ -181,73 +271,128 @@ func (u *StdConn) ListenOut(r EncReader) error {
|
||||
return err
|
||||
}
|
||||
|
||||
for i := 0; i < n; i++ {
|
||||
// Its ok to skip the ok check here, the slicing is the only error that can occur and it will panic
|
||||
if u.isV4 {
|
||||
ip, _ = netip.AddrFromSlice(names[i][4:8])
|
||||
} else {
|
||||
ip, _ = netip.AddrFromSlice(names[i][8:24])
|
||||
for i := range n {
|
||||
from := getFrom(names, i, u.isV4)
|
||||
payload := buffers[i][:msgs[i].Len]
|
||||
|
||||
segSize := 0
|
||||
if cmsgSpace > 0 {
|
||||
segSize = parseRecvCmsg(&msgs[i].Hdr)
|
||||
}
|
||||
r(netip.AddrPortFrom(ip.Unmap(), binary.BigEndian.Uint16(names[i][2:4])), buffers[i][:msgs[i].Len])
|
||||
|
||||
deliverSegments(r, from, payload, segSize)
|
||||
}
|
||||
|
||||
flush()
|
||||
}
|
||||
}
|
||||
|
||||
// deliverSegments hands a received superdatagram to r, splitting it back into pre-coalesce packets
|
||||
func deliverSegments(r EncReader, from netip.AddrPort, payload []byte, segSize int) {
|
||||
if segSize <= 0 || segSize >= len(payload) { //avoid bogus values
|
||||
r(from, payload[:len(payload):len(payload)])
|
||||
return
|
||||
}
|
||||
for off := 0; off < len(payload); off += segSize {
|
||||
end := off + segSize
|
||||
if end > len(payload) {
|
||||
end = len(payload)
|
||||
}
|
||||
r(from, payload[off:end:end])
|
||||
}
|
||||
}
|
||||
|
||||
// parseRecvCmsg walks the per-slot ancillary buffer and extracts the UDP_GRO
|
||||
// gso_size, or 0 when no UDP_GRO cmsg is present.
|
||||
func parseRecvCmsg(hdr *msghdr) (gso int) {
|
||||
controllen := int(hdr.Controllen)
|
||||
if controllen < unix.SizeofCmsghdr || hdr.Control == nil {
|
||||
return 0
|
||||
}
|
||||
ctrl := unsafe.Slice(hdr.Control, controllen)
|
||||
off := 0
|
||||
for off+unix.SizeofCmsghdr <= len(ctrl) {
|
||||
ch := (*unix.Cmsghdr)(unsafe.Pointer(&ctrl[off]))
|
||||
clen := int(ch.Len)
|
||||
// Compare against the remaining bytes rather than off+clen
|
||||
if clen < unix.SizeofCmsghdr || clen > len(ctrl)-off {
|
||||
return gso
|
||||
}
|
||||
dataOff := off + unix.CmsgLen(0)
|
||||
if ch.Level == unix.SOL_UDP && ch.Type == unix.UDP_GRO {
|
||||
if dataOff+udpGROCmsgPayload <= len(ctrl) {
|
||||
gso = int(int32(binary.NativeEndian.Uint32(ctrl[dataOff : dataOff+udpGROCmsgPayload])))
|
||||
}
|
||||
}
|
||||
// Advance by the aligned cmsg space.
|
||||
off += unix.CmsgSpace(clen - unix.CmsgLen(0))
|
||||
}
|
||||
return gso
|
||||
}
|
||||
|
||||
func (u *StdConn) WriteTo(b []byte, ip netip.AddrPort) error {
|
||||
if u.isV4 {
|
||||
return u.writeTo4(b, ip)
|
||||
}
|
||||
return u.writeTo6(b, ip)
|
||||
return sendto(u.sysFd, b, ip, u.isV4)
|
||||
}
|
||||
|
||||
func (u *StdConn) writeTo6(b []byte, ip netip.AddrPort) error {
|
||||
var rsa unix.RawSockaddrInet6
|
||||
rsa.Family = unix.AF_INET6
|
||||
rsa.Addr = ip.Addr().As16()
|
||||
binary.BigEndian.PutUint16((*[2]byte)(unsafe.Pointer(&rsa.Port))[:], ip.Port())
|
||||
|
||||
for {
|
||||
_, _, err := unix.Syscall6(
|
||||
unix.SYS_SENDTO,
|
||||
uintptr(u.sysFd),
|
||||
uintptr(unsafe.Pointer(&b[0])),
|
||||
uintptr(len(b)),
|
||||
uintptr(0),
|
||||
uintptr(unsafe.Pointer(&rsa)),
|
||||
uintptr(unix.SizeofSockaddrInet6),
|
||||
)
|
||||
if err != 0 {
|
||||
return &net.OpError{Op: "sendto", Err: err}
|
||||
}
|
||||
return nil
|
||||
func sendto(fd int, b []byte, addr netip.AddrPort, isV4 bool) error {
|
||||
var rsa [unix.SizeofSockaddrInet6]byte
|
||||
nlen, err := writeSockaddr(rsa[:], addr, isV4)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var base *byte
|
||||
if len(b) > 0 {
|
||||
base = &b[0]
|
||||
}
|
||||
_, _, errno := unix.Syscall6(
|
||||
unix.SYS_SENDTO,
|
||||
uintptr(fd),
|
||||
uintptr(unsafe.Pointer(base)),
|
||||
uintptr(len(b)),
|
||||
0,
|
||||
uintptr(unsafe.Pointer(&rsa[0])),
|
||||
uintptr(nlen),
|
||||
)
|
||||
if errno != 0 {
|
||||
return &net.OpError{Op: "sendto", Err: errno}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (u *StdConn) writeTo4(b []byte, ip netip.AddrPort) error {
|
||||
if !ip.Addr().Is4() {
|
||||
return ErrInvalidIPv6RemoteForSocket
|
||||
}
|
||||
// WriteBatch sends bufs via sendmmsg(2), coalescing same-destination runs into UDP-GSO superpackets when supported.
|
||||
// See batchWriter in udp_linux_writebatch.go for the mechanics.
|
||||
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
|
||||
return u.bw.WriteBatch(bufs, addrs)
|
||||
}
|
||||
|
||||
var rsa unix.RawSockaddrInet4
|
||||
rsa.Family = unix.AF_INET
|
||||
rsa.Addr = ip.Addr().As4()
|
||||
binary.BigEndian.PutUint16((*[2]byte)(unsafe.Pointer(&rsa.Port))[:], ip.Port())
|
||||
|
||||
for {
|
||||
_, _, err := unix.Syscall6(
|
||||
unix.SYS_SENDTO,
|
||||
uintptr(u.sysFd),
|
||||
uintptr(unsafe.Pointer(&b[0])),
|
||||
uintptr(len(b)),
|
||||
uintptr(0),
|
||||
uintptr(unsafe.Pointer(&rsa)),
|
||||
uintptr(unix.SizeofSockaddrInet4),
|
||||
)
|
||||
if err != 0 {
|
||||
return &net.OpError{Op: "sendto", Err: err}
|
||||
// writeSockaddr encodes addr into buf (which must be at least SizeofSockaddrInet6 bytes).
|
||||
// Returns the number of bytes used.
|
||||
// If isV4 is true and addr is not a v4 (or v4-in-v6) address, returns an error.
|
||||
func writeSockaddr(buf []byte, addr netip.AddrPort, isV4 bool) (int, error) {
|
||||
ap := addr.Addr().Unmap()
|
||||
if isV4 {
|
||||
if !ap.Is4() {
|
||||
return 0, ErrInvalidIPv6RemoteForSocket
|
||||
}
|
||||
return nil
|
||||
// struct sockaddr_in: { sa_family_t(2), in_port_t(2, BE), in_addr(4), zero(8) }
|
||||
// sa_family is host endian.
|
||||
binary.NativeEndian.PutUint16(buf[0:2], unix.AF_INET)
|
||||
binary.BigEndian.PutUint16(buf[2:4], addr.Port())
|
||||
ip4 := ap.As4()
|
||||
copy(buf[4:8], ip4[:])
|
||||
for j := 8; j < 16; j++ {
|
||||
buf[j] = 0
|
||||
}
|
||||
return unix.SizeofSockaddrInet4, nil
|
||||
}
|
||||
// struct sockaddr_in6: { sa_family_t(2), in_port_t(2, BE), flowinfo(4), in6_addr(16), scope_id(4) }
|
||||
binary.NativeEndian.PutUint16(buf[0:2], unix.AF_INET6)
|
||||
binary.BigEndian.PutUint16(buf[2:4], addr.Port())
|
||||
binary.NativeEndian.PutUint32(buf[4:8], 0)
|
||||
ip6 := addr.Addr().As16()
|
||||
copy(buf[8:24], ip6[:])
|
||||
binary.NativeEndian.PutUint32(buf[24:28], 0)
|
||||
return unix.SizeofSockaddrInet6, nil
|
||||
}
|
||||
|
||||
func (u *StdConn) ReloadConfig(c *config.C) {
|
||||
@@ -303,7 +448,7 @@ func (u *StdConn) getMemInfo(meminfo *[unix.SK_MEMINFO_VARS]uint32) error {
|
||||
|
||||
func (u *StdConn) Close() error {
|
||||
u.closed.Store(true)
|
||||
// Wake the reader parked in recvmmsg/recvmsg. shutdown(2) on an unconnected socket
|
||||
// Wake the reader parked in recvmmsg. shutdown(2) on an unconnected socket
|
||||
// returns ENOTCONN but still wakes it, so ignore the error.
|
||||
// The reader then sees closed and stops touching the fd, making the Close below safe.
|
||||
_ = unix.Shutdown(u.sysFd, unix.SHUT_RDWR)
|
||||
|
||||
+11
-18
@@ -30,25 +30,18 @@ type rawMessage struct {
|
||||
Len uint32
|
||||
}
|
||||
|
||||
func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) {
|
||||
msgs := make([]rawMessage, n)
|
||||
buffers := make([][]byte, n)
|
||||
names := make([][]byte, n)
|
||||
func setIovLen(v *iovec, n int) {
|
||||
v.Len = uint32(n)
|
||||
}
|
||||
|
||||
for i := range msgs {
|
||||
buffers[i] = make([]byte, MTU)
|
||||
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||
func setMsgIovlen(m *msghdr, n int) {
|
||||
m.Iovlen = uint32(n)
|
||||
}
|
||||
|
||||
vs := []iovec{
|
||||
{Base: &buffers[i][0], Len: uint32(len(buffers[i]))},
|
||||
}
|
||||
func setMsgControllen(m *msghdr, n int) {
|
||||
m.Controllen = uint32(n)
|
||||
}
|
||||
|
||||
msgs[i].Hdr.Iov = &vs[0]
|
||||
msgs[i].Hdr.Iovlen = uint32(len(vs))
|
||||
|
||||
msgs[i].Hdr.Name = &names[i][0]
|
||||
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
||||
}
|
||||
|
||||
return msgs, buffers, names
|
||||
func setCmsgLen(h *unix.Cmsghdr, n int) {
|
||||
h.Len = uint32(n)
|
||||
}
|
||||
|
||||
+11
-18
@@ -33,25 +33,18 @@ type rawMessage struct {
|
||||
Pad0 [4]byte
|
||||
}
|
||||
|
||||
func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) {
|
||||
msgs := make([]rawMessage, n)
|
||||
buffers := make([][]byte, n)
|
||||
names := make([][]byte, n)
|
||||
func setIovLen(v *iovec, n int) {
|
||||
v.Len = uint64(n)
|
||||
}
|
||||
|
||||
for i := range msgs {
|
||||
buffers[i] = make([]byte, MTU)
|
||||
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||
func setMsgIovlen(m *msghdr, n int) {
|
||||
m.Iovlen = uint64(n)
|
||||
}
|
||||
|
||||
vs := []iovec{
|
||||
{Base: &buffers[i][0], Len: uint64(len(buffers[i]))},
|
||||
}
|
||||
func setMsgControllen(m *msghdr, n int) {
|
||||
m.Controllen = uint64(n)
|
||||
}
|
||||
|
||||
msgs[i].Hdr.Iov = &vs[0]
|
||||
msgs[i].Hdr.Iovlen = uint64(len(vs))
|
||||
|
||||
msgs[i].Hdr.Name = &names[i][0]
|
||||
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
||||
}
|
||||
|
||||
return msgs, buffers, names
|
||||
func setCmsgLen(h *unix.Cmsghdr, n int) {
|
||||
h.Len = uint64(n)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,708 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"testing"
|
||||
"time"
|
||||
"unsafe"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// TestGSOMaxSegmentsKernelGate pins the corrected kernel-version gate: the
|
||||
// 128-segment cap (127 usable) only lands in Linux v6.9 (commit 1382e3b6a350),
|
||||
// not 5.5. Everything older stays at the conservative 63.
|
||||
func TestGSOMaxSegmentsKernelGate(t *testing.T) {
|
||||
cases := []struct {
|
||||
release string
|
||||
want int
|
||||
}{
|
||||
{"5.4.0", 63},
|
||||
{"5.5.0-generic", 63}, // the old bug bumped here — it must not now
|
||||
{"5.15.0", 63},
|
||||
{"6.1.0", 63},
|
||||
{"6.8.0-generic", 63},
|
||||
{"6.9.0", 127},
|
||||
{"6.10.1-arch1-1", 127},
|
||||
{"7.0.5-arch1-1", 127},
|
||||
{"garbage", 63},
|
||||
{"", 63},
|
||||
}
|
||||
for _, c := range cases {
|
||||
if got := gsoMaxSegments(c.release); got != c.want {
|
||||
t.Errorf("gsoMaxSegments(%q) = %d, want %d", c.release, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// buildCmsg lays out a single ancillary cmsg (header + data) in a fresh buffer
|
||||
// the way the kernel would deliver it, so parseRecvCmsg can be exercised
|
||||
// without a live socket.
|
||||
func buildCmsg(level, typ int32, data []byte) []byte {
|
||||
buf := make([]byte, unix.CmsgSpace(len(data)))
|
||||
h := (*unix.Cmsghdr)(unsafe.Pointer(&buf[0]))
|
||||
h.Level = level
|
||||
h.Type = typ
|
||||
setCmsgLen(h, unix.CmsgLen(len(data)))
|
||||
copy(buf[unix.CmsgLen(0):], data)
|
||||
return buf
|
||||
}
|
||||
|
||||
func testLogger() *slog.Logger {
|
||||
return slog.New(slog.DiscardHandler)
|
||||
}
|
||||
|
||||
// TestWriteBatchBadFamilyDeliversOthers is the H3 regression: a batch that
|
||||
// contains one destination the socket can't reach (an IPv6 remote on a
|
||||
// v4-bound socket) must still deliver every other packet. Before the fix the
|
||||
// writeSockaddr error returned early and dropped the whole chunk.
|
||||
func TestWriteBatchBadFamilyDeliversOthers(t *testing.T) {
|
||||
rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
|
||||
if err != nil {
|
||||
t.Skipf("cannot open v4 receiver (sandbox?): %v", err)
|
||||
}
|
||||
defer rx.Close()
|
||||
rxPort := rx.LocalAddr().(*net.UDPAddr).Port
|
||||
|
||||
// Bind a *non-wildcard* v4 address so Go gives us a genuine AF_INET
|
||||
// socket. A wildcard v4 bind (0.0.0.0) via network "udp" comes up as a
|
||||
// dual-stack AF_INET6 socket on Linux, for which a v6 dest is not a bad
|
||||
// family — which would defeat the point of this test.
|
||||
udpSettings := Settings{
|
||||
Listen: netip.MustParseAddrPort("127.0.0.1:0"),
|
||||
Multi: false,
|
||||
Batch: 1,
|
||||
Offloads: false,
|
||||
}
|
||||
c, err := NewListener(testLogger(), udpSettings)
|
||||
if err != nil {
|
||||
t.Skipf("cannot open v4 sender (sandbox?): %v", err)
|
||||
}
|
||||
defer c.Close()
|
||||
sender := c.(*StdConn)
|
||||
if !sender.isV4 {
|
||||
t.Fatalf("expected a v4-bound sender socket, got isV4=false")
|
||||
}
|
||||
|
||||
good := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), uint16(rxPort))
|
||||
bad := netip.MustParseAddrPort("[2001:db8::1]:9999") // genuine v6, unreachable on v4 socket
|
||||
|
||||
bufs := [][]byte{[]byte("AAA"), []byte("BBB"), []byte("CCC")}
|
||||
addrs := []netip.AddrPort{good, bad, good}
|
||||
|
||||
n, err := sender.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteBatch returned error, want nil (bad dest should be isolated): %v", err)
|
||||
}
|
||||
if n != 2 {
|
||||
t.Errorf("WriteBatch wrote %d packets, want 2 of 3 (the bad-family dest is the only casualty)", n)
|
||||
}
|
||||
|
||||
got := map[string]bool{}
|
||||
rx.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
buf := make([]byte, 64)
|
||||
for i := 0; i < 2; i++ {
|
||||
n, _, rerr := rx.ReadFromUDPAddrPort(buf)
|
||||
if rerr != nil {
|
||||
t.Fatalf("expected 2 delivered packets, read #%d failed: %v", i+1, rerr)
|
||||
}
|
||||
got[string(buf[:n])] = true
|
||||
}
|
||||
if !got["AAA"] || !got["CCC"] {
|
||||
t.Errorf("delivered set = %v, want AAA and CCC both present", got)
|
||||
}
|
||||
if got["BBB"] {
|
||||
t.Errorf("the bad-family packet BBB was somehow delivered")
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteBatchUnreachableDestDeliversOthers is the kernel-rejection twin of
|
||||
// TestWriteBatchBadFamilyDeliversOthers. A destination the kernel refuses outright (240.0.0.0/4 is reserved, so
|
||||
// the send returns EINVAL) fails its sendmmsg entry; WriteBatch must drop only that entry and still deliver
|
||||
// every other packet rather than abandoning the batch at the first failure.
|
||||
func TestWriteBatchUnreachableDestDeliversOthers(t *testing.T) {
|
||||
rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
|
||||
if err != nil {
|
||||
t.Skipf("cannot open v4 receiver (sandbox?): %v", err)
|
||||
}
|
||||
defer rx.Close()
|
||||
rxPort := rx.LocalAddr().(*net.UDPAddr).Port
|
||||
|
||||
udpSettings := Settings{
|
||||
Listen: netip.MustParseAddrPort("127.0.0.1:0"),
|
||||
Multi: false,
|
||||
Batch: 1,
|
||||
Offloads: false,
|
||||
}
|
||||
c, err := NewListener(testLogger(), udpSettings)
|
||||
if err != nil {
|
||||
t.Skipf("cannot open v4 sender (sandbox?): %v", err)
|
||||
}
|
||||
defer c.Close()
|
||||
sender := c.(*StdConn)
|
||||
|
||||
good := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), uint16(rxPort))
|
||||
bad := netip.MustParseAddrPort("240.0.0.1:9999") // reserved space, the kernel refuses it
|
||||
|
||||
bufs := [][]byte{[]byte("P0"), []byte("P1"), []byte("BAD"), []byte("P3"), []byte("P4")}
|
||||
addrs := []netip.AddrPort{good, good, bad, good, good}
|
||||
|
||||
// The bad destination is reported, but only after every other packet has been attempted.
|
||||
if _, err := sender.WriteBatch(bufs, addrs); err == nil {
|
||||
t.Log("WriteBatch returned nil; kernel accepted the reserved address, delivery assertions still apply")
|
||||
}
|
||||
|
||||
got := map[string]bool{}
|
||||
rx.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
buf := make([]byte, 64)
|
||||
for i := 0; i < 4; i++ {
|
||||
n, _, rerr := rx.ReadFromUDPAddrPort(buf)
|
||||
if rerr != nil {
|
||||
t.Fatalf("expected 4 delivered packets, read #%d failed: %v (got so far: %v)", i+1, rerr, got)
|
||||
}
|
||||
got[string(buf[:n])] = true
|
||||
}
|
||||
for _, want := range []string{"P0", "P1", "P3", "P4"} {
|
||||
if !got[want] {
|
||||
t.Errorf("packet %s was not delivered; delivered set = %v", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseRecvCmsgCorruptLenNoPanic: a cmsg Len near max-int used to wrap
|
||||
// off+clen negative, slip past the bounds check, and drive the walk offset
|
||||
// negative -- a panic on the next ctrl[off]. The guard must compare Len
|
||||
// against the remaining bytes instead. Also pins the plain truncated-Len
|
||||
// cases (too small, larger than the buffer) to a clean early return.
|
||||
func TestParseRecvCmsgCorruptLenNoPanic(t *testing.T) {
|
||||
// First cmsg: a valid empty one so the walk advances past off=0
|
||||
// (off+clen can't overflow while off is still zero).
|
||||
valid := buildCmsg(int32(unix.SOL_UDP), int32(unix.UDP_GRO), make([]byte, 4))
|
||||
|
||||
corrupt := func(lenVal int) []byte {
|
||||
buf := make([]byte, len(valid)+unix.CmsgSpace(4))
|
||||
copy(buf, valid)
|
||||
h := (*unix.Cmsghdr)(unsafe.Pointer(&buf[len(valid)]))
|
||||
h.Level = int32(unix.IPPROTO_IP)
|
||||
h.Type = int32(unix.IP_TOS)
|
||||
setCmsgLen(h, lenVal)
|
||||
return buf
|
||||
}
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
ctrl []byte
|
||||
}{
|
||||
{"len_near_max_int", corrupt(int(^uint(0)>>1) - 8)},
|
||||
{"len_too_small", corrupt(unix.SizeofCmsghdr - 1)},
|
||||
{"len_past_buffer", corrupt(1 << 20)},
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
hdr := &msghdr{Control: &c.ctrl[0]}
|
||||
setMsgControllen(hdr, len(c.ctrl))
|
||||
gso := parseRecvCmsg(hdr)
|
||||
// The valid leading UDP_GRO cmsg (payload 0) must still parse;
|
||||
// the corrupt trailer just ends the walk.
|
||||
if gso != 0 {
|
||||
t.Errorf("parseRecvCmsg = %d, want 0", gso)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestDeliverSegments pins the GRO RX splitting: a kernel-coalesced buffer
|
||||
// must come back out as the exact pre-coalesce packets -- every boundary
|
||||
// error here shreds encrypted packets and every decrypt downstream fails.
|
||||
func TestDeliverSegments(t *testing.T) {
|
||||
from := netip.MustParseAddrPort("192.0.2.1:4242")
|
||||
// Spare backing capacity mimics the recvmmsg row a real payload sits in;
|
||||
// the cap checks below prove none of it leaks to a delivered segment.
|
||||
pay := func(n int) []byte {
|
||||
b := make([]byte, n, n+512)
|
||||
for i := range b {
|
||||
b[i] = byte(i)
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
payload []byte
|
||||
segSize int
|
||||
wantLens []int
|
||||
}{
|
||||
{"no-gro", pay(1400), 0, []int{1400}},
|
||||
{"negative-segsize", pay(1400), -5, []int{1400}},
|
||||
{"segsize-equals-payload", pay(1400), 1400, []int{1400}},
|
||||
{"segsize-past-payload", pay(1400), 2000, []int{1400}},
|
||||
{"even-split", pay(4200), 1400, []int{1400, 1400, 1400}},
|
||||
{"short-tail", pay(3000), 1400, []int{1400, 1400, 200}},
|
||||
{"single-byte-tail", pay(2801), 1400, []int{1400, 1400, 1}},
|
||||
{"segsize-one", pay(3), 1, []int{1, 1, 1}},
|
||||
{"empty-payload", pay(0), 1400, []int{0}},
|
||||
{"max-coalesce", pay(65500), 1372, nil}, // lens derived below
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
wantLens := c.wantLens
|
||||
if wantLens == nil {
|
||||
for rem := len(c.payload); rem > 0; rem -= c.segSize {
|
||||
wantLens = append(wantLens, min(c.segSize, rem))
|
||||
}
|
||||
}
|
||||
|
||||
var got [][]byte
|
||||
deliverSegments(func(a netip.AddrPort, seg []byte) {
|
||||
if a != from {
|
||||
t.Errorf("from = %v, want %v", a, from)
|
||||
}
|
||||
got = append(got, seg)
|
||||
}, from, c.payload, c.segSize)
|
||||
|
||||
if len(got) != len(wantLens) {
|
||||
t.Fatalf("delivered %d segments, want %d", len(got), len(wantLens))
|
||||
}
|
||||
// Segments must tile the payload in order with no gap, overlap,
|
||||
// or copy: each must alias the payload at the right offset.
|
||||
off := 0
|
||||
for i, seg := range got {
|
||||
if len(seg) != wantLens[i] {
|
||||
t.Fatalf("segment %d len=%d want %d", i, len(seg), wantLens[i])
|
||||
}
|
||||
if cap(seg) != len(seg) {
|
||||
// EncReader contract: an append into spare capacity would
|
||||
// scribble into the next segment of the shared row.
|
||||
t.Errorf("segment %d cap=%d, want %d (capacity must not reach into the row)", i, cap(seg), len(seg))
|
||||
}
|
||||
if len(seg) > 0 && &seg[0] != &c.payload[off] {
|
||||
t.Errorf("segment %d does not alias payload at offset %d", i, off)
|
||||
}
|
||||
off += len(seg)
|
||||
}
|
||||
if off != len(c.payload) {
|
||||
t.Errorf("segments cover %d bytes, payload has %d", off, len(c.payload))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// newRewindTestWriter builds a batchWriter with no socket: GSO planning on,
|
||||
// sendFn left for the test to script. fd is invalid on purpose -- any path
|
||||
// that actually hits the kernel fails loudly.
|
||||
func newRewindTestWriter() *batchWriter {
|
||||
w := &batchWriter{fd: -1, isV4: true, l: testLogger()}
|
||||
// gsoSupported must be set before prepareWriteMessages: the cmsg slab is
|
||||
// only allocated when GSO is already known to be supported.
|
||||
w.gsoSupported = true
|
||||
w.maxGSOSegments = 63
|
||||
w.prepareWriteMessages(MaxWriteBatch, true)
|
||||
return w
|
||||
}
|
||||
|
||||
// capturePrepared decodes n prepared mmsghdr entries beginning at start
|
||||
// straight from their iovecs -- ground truth, deliberately not the entryEnd
|
||||
// bookkeeping the resume logic itself relies on. Returns one []byte per
|
||||
// packed packet, in entry order.
|
||||
func capturePrepared(w *batchWriter, start, n int) [][]byte {
|
||||
var out [][]byte
|
||||
for e := start; e < start+n; e++ {
|
||||
hdr := &w.msgs[e].Hdr
|
||||
iovs := unsafe.Slice(hdr.Iov, int(hdr.Iovlen))
|
||||
for _, iov := range iovs {
|
||||
b := make([]byte, int(iov.Len))
|
||||
if iov.Len > 0 {
|
||||
copy(b, unsafe.Slice(iov.Base, int(iov.Len)))
|
||||
}
|
||||
out = append(out, b)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// TestWriteBatchPartialSendRewind drives WriteBatch through scripted
|
||||
// partial sendmmsg results and asserts the rewind resumes exactly where
|
||||
// the kernel stopped: every packet on the wire exactly once, in order,
|
||||
// no duplicate, no loss. This is the hairiest logic in the write path
|
||||
// and a rewind bug means silent packet duplication or loss under EAGAIN-
|
||||
// style backpressure.
|
||||
func TestWriteBatchPartialSendRewind(t *testing.T) {
|
||||
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
|
||||
|
||||
mkBuf := func(tag byte, n int) []byte {
|
||||
b := make([]byte, n)
|
||||
for i := range b {
|
||||
b[i] = tag
|
||||
}
|
||||
b[0] = tag // tag identifies the packet uniquely below
|
||||
return b
|
||||
}
|
||||
|
||||
// Mixed shape: a 3-packet GSO run to A, a lone short packet to A (run
|
||||
// tail), then two to B. The planner packs this as multiple entries with
|
||||
// multi-iovec runs, which is what makes the rewind arithmetic hairy.
|
||||
bufs := [][]byte{
|
||||
mkBuf(1, 1200), mkBuf(2, 1200), mkBuf(3, 1200), // run to A
|
||||
mkBuf(4, 600), // short tail to A
|
||||
mkBuf(5, 900), mkBuf(6, 900), // run to B
|
||||
}
|
||||
addrs := []netip.AddrPort{dstA, dstA, dstA, dstA, dstB, dstB}
|
||||
|
||||
scripts := [][]int{
|
||||
{99}, // accept everything first call
|
||||
{1, 99}, // one entry per call, then the rest
|
||||
{1, 1, 1, 99}, // strictly one entry per call
|
||||
{2, 99}, // two entries, then the rest
|
||||
}
|
||||
for si, script := range scripts {
|
||||
t.Run(fmt.Sprintf("script_%d", si), func(t *testing.T) {
|
||||
w := newRewindTestWriter()
|
||||
var wire [][]byte
|
||||
call := 0
|
||||
w.sendFn = func(start, n int) (int, error) {
|
||||
accept := n
|
||||
if call < len(script) && script[call] < n {
|
||||
accept = script[call]
|
||||
}
|
||||
call++
|
||||
wire = append(wire, capturePrepared(w, start, accept)...)
|
||||
return accept, nil
|
||||
}
|
||||
|
||||
written, err := w.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteBatch: %v", err)
|
||||
}
|
||||
if written != len(bufs) {
|
||||
t.Errorf("written = %d, want %d", written, len(bufs))
|
||||
}
|
||||
if len(wire) != len(bufs) {
|
||||
t.Fatalf("wire got %d packets, want %d (dup or loss in rewind)", len(wire), len(bufs))
|
||||
}
|
||||
for i, b := range wire {
|
||||
if len(b) != len(bufs[i]) || b[0] != bufs[i][0] {
|
||||
t.Errorf("wire[%d] = tag %d len %d, want tag %d len %d (reorder/dup)",
|
||||
i, b[0], len(b), bufs[i][0], len(bufs[i]))
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteBatchSkipUnroutableRunAccounting: an unroutable destination mid-
|
||||
// batch is skipped without committing an entry, leaving a hole in the bufs
|
||||
// index space. The written count must tally packets per sent entry -- the
|
||||
// index span would count the hole -- across both full and partial sendmmsg
|
||||
// success.
|
||||
func TestWriteBatchSkipUnroutableRunAccounting(t *testing.T) {
|
||||
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
|
||||
bad := netip.MustParseAddrPort("[2001:db8::1]:9999") // v6 dest, v4 writer
|
||||
|
||||
mk := func(tag byte, n int) []byte {
|
||||
b := make([]byte, n)
|
||||
b[0] = tag
|
||||
return b
|
||||
}
|
||||
bufs := [][]byte{mk(1, 1200), mk(2, 1200), mk(3, 500), mk(4, 900), mk(5, 900)}
|
||||
addrs := []netip.AddrPort{dstA, dstA, bad, dstB, dstB}
|
||||
|
||||
for si, script := range [][]int{{99}, {1, 99}} {
|
||||
t.Run(fmt.Sprintf("script_%d", si), func(t *testing.T) {
|
||||
w := newRewindTestWriter()
|
||||
var wire [][]byte
|
||||
call := 0
|
||||
w.sendFn = func(start, n int) (int, error) {
|
||||
accept := n
|
||||
if call < len(script) && script[call] < n {
|
||||
accept = script[call]
|
||||
}
|
||||
call++
|
||||
wire = append(wire, capturePrepared(w, start, accept)...)
|
||||
return accept, nil
|
||||
}
|
||||
|
||||
written, err := w.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteBatch: %v", err)
|
||||
}
|
||||
if written != 4 {
|
||||
t.Errorf("written = %d, want 4 (the unroutable run is the only casualty)", written)
|
||||
}
|
||||
wantTags := []byte{1, 2, 4, 5}
|
||||
if len(wire) != len(wantTags) {
|
||||
t.Fatalf("wire got %d packets, want %d (dup or loss around the skip)", len(wire), len(wantTags))
|
||||
}
|
||||
for i, b := range wire {
|
||||
if b[0] != wantTags[i] {
|
||||
t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i])
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteBatchMidChunkRejectResumes: after a partial success, a zero-sent
|
||||
// error on the FIRST REMAINING entry (done > 0) must drop only that entry's
|
||||
// run and resume the rest of the chunk in place -- no repacking, no packets
|
||||
// lost from entries before or after the rejected one.
|
||||
func TestWriteBatchMidChunkRejectResumes(t *testing.T) {
|
||||
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
|
||||
dstC := netip.MustParseAddrPort("127.0.0.3:4242")
|
||||
|
||||
mk := func(tag byte, n int) []byte {
|
||||
b := make([]byte, n)
|
||||
b[0] = tag
|
||||
return b
|
||||
}
|
||||
// Three entries: a 2-packet GSO run to A, a 2-packet run to B, one to C.
|
||||
bufs := [][]byte{mk(1, 1200), mk(2, 1200), mk(3, 900), mk(4, 900), mk(5, 600)}
|
||||
addrs := []netip.AddrPort{dstA, dstA, dstB, dstB, dstC}
|
||||
|
||||
w := newRewindTestWriter()
|
||||
var wire [][]byte
|
||||
var starts []int
|
||||
call := 0
|
||||
w.sendFn = func(start, n int) (int, error) {
|
||||
starts = append(starts, start)
|
||||
call++
|
||||
switch call {
|
||||
case 1: // accept only entry 0 (the run to A)
|
||||
wire = append(wire, capturePrepared(w, start, 1)...)
|
||||
return 1, nil
|
||||
case 2: // reject entry 1 (the run to B) outright
|
||||
return -1, &net.OpError{Op: "sendmmsg", Err: unix.EPERM}
|
||||
default: // accept the rest
|
||||
wire = append(wire, capturePrepared(w, start, n)...)
|
||||
return n, nil
|
||||
}
|
||||
}
|
||||
|
||||
written, err := w.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteBatch: %v", err)
|
||||
}
|
||||
if written != 3 {
|
||||
t.Errorf("written = %d, want 3 (B's rejected run is the only casualty)", written)
|
||||
}
|
||||
wantTags := []byte{1, 2, 5}
|
||||
if len(wire) != len(wantTags) {
|
||||
t.Fatalf("wire got %d packets, want %d (dup or loss around the mid-chunk reject)", len(wire), len(wantTags))
|
||||
}
|
||||
for i, b := range wire {
|
||||
if b[0] != wantTags[i] {
|
||||
t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i])
|
||||
}
|
||||
}
|
||||
// The resume must reuse the prepared entries: same chunk, advancing
|
||||
// start offsets, no repack (which would restart at 0 with fresh entries).
|
||||
if want := []int{0, 1, 2}; !slices.Equal(starts, want) {
|
||||
t.Errorf("sendFn start offsets = %v, want %v", starts, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteBatchMidChunkEIODisablesGSOWithoutDup: an EIO on a GSO entry
|
||||
// after earlier entries in the chunk already went out must replay ONLY from
|
||||
// the failed run (replanned as single-packet entries) -- the already-sent
|
||||
// entries must not be duplicated.
|
||||
func TestWriteBatchMidChunkEIODisablesGSOWithoutDup(t *testing.T) {
|
||||
dstA := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||
dstB := netip.MustParseAddrPort("127.0.0.2:4242")
|
||||
|
||||
mk := func(tag byte, n int) []byte {
|
||||
b := make([]byte, n)
|
||||
b[0] = tag
|
||||
return b
|
||||
}
|
||||
// Entry 0: single packet to A. Entry 1: 2-packet GSO run to B.
|
||||
bufs := [][]byte{mk(1, 600), mk(2, 1200), mk(3, 1200)}
|
||||
addrs := []netip.AddrPort{dstA, dstB, dstB}
|
||||
|
||||
w := newRewindTestWriter()
|
||||
var wire [][]byte
|
||||
call := 0
|
||||
w.sendFn = func(start, n int) (int, error) {
|
||||
call++
|
||||
switch call {
|
||||
case 1: // accept entry 0 only
|
||||
wire = append(wire, capturePrepared(w, start, 1)...)
|
||||
return 1, nil
|
||||
case 2: // EIO on the GSO run to B
|
||||
return -1, &net.OpError{Op: "sendmmsg", Err: unix.EIO}
|
||||
default: // replanned single-packet replay
|
||||
wire = append(wire, capturePrepared(w, start, n)...)
|
||||
return n, nil
|
||||
}
|
||||
}
|
||||
|
||||
written, err := w.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteBatch: %v", err)
|
||||
}
|
||||
if w.gsoSupported {
|
||||
t.Error("gsoSupported still true after EIO on a GSO entry")
|
||||
}
|
||||
if written != len(bufs) {
|
||||
t.Errorf("written = %d, want %d", written, len(bufs))
|
||||
}
|
||||
wantTags := []byte{1, 2, 3}
|
||||
if len(wire) != len(wantTags) {
|
||||
t.Fatalf("wire got %d packets, want %d (packet 1 duplicated, or B's run lost)", len(wire), len(wantTags))
|
||||
}
|
||||
for i, b := range wire {
|
||||
if b[0] != wantTags[i] {
|
||||
t.Errorf("wire[%d] tag = %d, want %d", i, b[0], wantTags[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteBatchZeroProgress: sent == 0 with no error must abort with an
|
||||
// error rather than spin forever replaying the same chunk.
|
||||
func TestWriteBatchZeroProgress(t *testing.T) {
|
||||
w := newRewindTestWriter()
|
||||
w.sendFn = func(start, n int) (int, error) { return 0, nil }
|
||||
bufs := [][]byte{make([]byte, 100)}
|
||||
addrs := []netip.AddrPort{netip.MustParseAddrPort("127.0.0.1:4242")}
|
||||
if _, err := w.WriteBatch(bufs, addrs); err == nil {
|
||||
t.Fatal("WriteBatch = nil error on zero progress, want error")
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteBatchEIODisablesGSOAndReplays pins the runtime GSO give-up: a
|
||||
// sendmmsg rejected with EIO on a GSO superpacket entry must clear
|
||||
// gsoSupported and replay the same packets as per-packet entries through
|
||||
// sendmmsg (keeping batching), not fall back to per-packet sendto.
|
||||
func TestWriteBatchEIODisablesGSOAndReplays(t *testing.T) {
|
||||
dst := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||
bufs := [][]byte{make([]byte, 1200), make([]byte, 1200), make([]byte, 1200)}
|
||||
addrs := []netip.AddrPort{dst, dst, dst}
|
||||
|
||||
w := newRewindTestWriter()
|
||||
var entryCounts []int
|
||||
call := 0
|
||||
w.sendFn = func(start, n int) (int, error) {
|
||||
entryCounts = append(entryCounts, n)
|
||||
call++
|
||||
if call == 1 {
|
||||
return -1, &net.OpError{Op: "sendmmsg", Err: unix.EIO}
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
written, err := w.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteBatch: %v", err)
|
||||
}
|
||||
if w.gsoSupported {
|
||||
t.Error("gsoSupported still true after EIO on a GSO entry")
|
||||
}
|
||||
if written != len(bufs) {
|
||||
t.Errorf("written = %d, want %d", written, len(bufs))
|
||||
}
|
||||
// First call: one GSO entry carrying the whole run. Replay: one entry
|
||||
// per packet, still via sendmmsg.
|
||||
want := []int{1, 3}
|
||||
if len(entryCounts) != len(want) || entryCounts[0] != want[0] || entryCounts[1] != want[1] {
|
||||
t.Errorf("sendmmsg entry counts = %v, want %v", entryCounts, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestGSOEngagesOnLoopback is the offload smoke test: real sockets, real
|
||||
// UDP_SEGMENT cmsg, real kernel segmentation over loopback. It asserts
|
||||
// both that GSO *engaged* (the whole batch left in a single sendmmsg
|
||||
// entry -- a silent fallback to per-packet entries fails the test) and
|
||||
// that the kernel carved the superpacket back into the exact original
|
||||
// datagrams on the receive side. Runs in CI (make test on ubuntu-latest),
|
||||
// which is what guards against the offload path silently degrading.
|
||||
func TestGSOEngagesOnLoopback(t *testing.T) {
|
||||
rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
|
||||
if err != nil {
|
||||
t.Fatalf("listen rx: %v", err)
|
||||
}
|
||||
defer rx.Close()
|
||||
dst := rx.LocalAddr().(*net.UDPAddr).AddrPort()
|
||||
|
||||
udpSettings := Settings{
|
||||
Listen: netip.MustParseAddrPort("127.0.0.1:0"),
|
||||
Multi: false,
|
||||
Batch: 8,
|
||||
Offloads: true,
|
||||
}
|
||||
uc, err := NewListener(testLogger(), udpSettings)
|
||||
if err != nil {
|
||||
t.Fatalf("NewListener: %v", err)
|
||||
}
|
||||
sc := uc.(*StdConn)
|
||||
defer sc.Close()
|
||||
|
||||
if !sc.bw.gsoSupported {
|
||||
var un unix.Utsname
|
||||
_ = unix.Uname(&un)
|
||||
release := string(un.Release[:])
|
||||
if major, minor := parseRelease(release); major > 4 || (major == 4 && minor >= 18) {
|
||||
t.Fatalf("kernel %q supports UDP_SEGMENT but the GSO probe failed", release)
|
||||
}
|
||||
t.Skipf("kernel %q predates UDP_SEGMENT (4.18)", release)
|
||||
}
|
||||
|
||||
// Spy on the real syscall to count entries per sendmmsg without
|
||||
// changing what hits the kernel.
|
||||
var entryCounts []int
|
||||
real := sc.bw.sendFn
|
||||
sc.bw.sendFn = func(start, n int) (int, error) {
|
||||
entryCounts = append(entryCounts, n)
|
||||
return real(start, n)
|
||||
}
|
||||
|
||||
const numPkts = 8
|
||||
const pktLen = 1200
|
||||
bufs := make([][]byte, numPkts)
|
||||
addrs := make([]netip.AddrPort, numPkts)
|
||||
for i := range bufs {
|
||||
bufs[i] = make([]byte, pktLen)
|
||||
for j := range bufs[i] {
|
||||
bufs[i][j] = byte(i)
|
||||
}
|
||||
addrs[i] = dst
|
||||
}
|
||||
|
||||
written, err := sc.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteBatch: %v", err)
|
||||
}
|
||||
if written != numPkts {
|
||||
t.Fatalf("written = %d, want %d", written, numPkts)
|
||||
}
|
||||
// GSO engaged means the run went out as ONE sendmmsg entry carrying a
|
||||
// UDP_SEGMENT superpacket. Per-packet entries mean it silently fell
|
||||
// back -- exactly the regression this test exists to catch.
|
||||
if len(entryCounts) != 1 || entryCounts[0] != 1 {
|
||||
t.Fatalf("sendmmsg entry counts = %v, want [1]: GSO did not engage", entryCounts)
|
||||
}
|
||||
|
||||
// The kernel must deliver the original datagram boundaries and bytes.
|
||||
_ = rx.SetReadDeadline(time.Now().Add(5 * time.Second))
|
||||
got := make([]byte, pktLen+1)
|
||||
for i := 0; i < numPkts; i++ {
|
||||
n, _, err := rx.ReadFromUDP(got)
|
||||
if err != nil {
|
||||
t.Fatalf("rx read %d: %v", i, err)
|
||||
}
|
||||
if n != pktLen {
|
||||
t.Fatalf("rx read %d: len=%d want %d (kernel segmented at wrong boundary)", i, n, pktLen)
|
||||
}
|
||||
for j := 0; j < n; j++ {
|
||||
if got[j] != byte(i) {
|
||||
t.Fatalf("rx read %d: byte %d = %#x, want %#x", i, j, got[j], byte(i))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,277 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// These tests pin the listen.udp_offloads=false behavior: no GSO/GRO probes,
|
||||
// no cmsg scratch, and — critically — a still-functional send/receive path.
|
||||
// The sockaddr name buffers are needed for every sendmmsg entry whether or
|
||||
// not offloads are on, so prepareWriteMessages must allocate them even when
|
||||
// it skips the cmsg slab (a nil name buffer panics in writeSockaddr on the
|
||||
// first WriteBatch).
|
||||
|
||||
// TestPrepareWriteMessagesAlwaysAllocatesNames covers all four
|
||||
// (offloadsEnabled, gsoSupported) combinations: the sockaddr name buffers
|
||||
// must exist in every one, and the cmsg slab only when both are true.
|
||||
// gsoSupported=false with offloads enabled is the old-kernel path where the
|
||||
// UDP_SEGMENT probe fails — not just a config choice.
|
||||
func TestPrepareWriteMessagesAlwaysAllocatesNames(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
offloads bool
|
||||
gso bool
|
||||
}{
|
||||
{"offloads-off", false, false},
|
||||
{"offloads-on-probe-failed", true, false},
|
||||
{"offloads-off-gso-flag-set", false, true},
|
||||
{"offloads-on", true, true},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
w := &batchWriter{fd: -1, isV4: true, l: testLogger()}
|
||||
w.gsoSupported = tc.gso
|
||||
w.prepareWriteMessages(MaxWriteBatch, tc.offloads)
|
||||
|
||||
for i := range w.msgs {
|
||||
if len(w.names[i]) != unix.SizeofSockaddrInet6 {
|
||||
t.Fatalf("names[%d] len=%d, want %d", i, len(w.names[i]), unix.SizeofSockaddrInet6)
|
||||
}
|
||||
if w.msgs[i].Hdr.Name == nil {
|
||||
t.Fatalf("msgs[%d].Hdr.Name is nil", i)
|
||||
}
|
||||
}
|
||||
|
||||
wantCmsg := tc.offloads && tc.gso
|
||||
if (w.cmsg != nil) != wantCmsg {
|
||||
t.Errorf("cmsg allocated = %v, want %v", w.cmsg != nil, wantCmsg)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteBatchOffloadsDisabledScripted drives WriteBatch through a
|
||||
// batchWriter built with offloads disabled and a scripted sendFn: every
|
||||
// packet must become its own sendmmsg entry (no GSO coalescing to plan),
|
||||
// packed correctly despite the missing cmsg slab.
|
||||
func TestWriteBatchOffloadsDisabledScripted(t *testing.T) {
|
||||
w := &batchWriter{fd: -1, isV4: true, l: testLogger()}
|
||||
w.prepareWriteMessages(MaxWriteBatch, false)
|
||||
|
||||
var entryCounts []int
|
||||
w.sendFn = func(start, n int) (int, error) {
|
||||
entryCounts = append(entryCounts, n)
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// Same destination, equal sizes: prime coalescing bait that must not
|
||||
// coalesce with offloads off.
|
||||
dst := netip.MustParseAddrPort("127.0.0.1:4242")
|
||||
const numPkts = 4
|
||||
bufs := make([][]byte, numPkts)
|
||||
addrs := make([]netip.AddrPort, numPkts)
|
||||
for i := range bufs {
|
||||
bufs[i] = make([]byte, 1200)
|
||||
for j := range bufs[i] {
|
||||
bufs[i][j] = byte(i)
|
||||
}
|
||||
addrs[i] = dst
|
||||
}
|
||||
|
||||
written, err := w.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteBatch: %v", err)
|
||||
}
|
||||
if written != numPkts {
|
||||
t.Errorf("written = %d, want %d", written, numPkts)
|
||||
}
|
||||
if len(entryCounts) != 1 || entryCounts[0] != numPkts {
|
||||
t.Errorf("sendmmsg entry counts = %v, want [%d]: packets must be one entry each", entryCounts, numPkts)
|
||||
}
|
||||
|
||||
// Each prepared entry must carry exactly its own packet's bytes.
|
||||
got := capturePrepared(w, 0, numPkts)
|
||||
if len(got) != numPkts {
|
||||
t.Fatalf("prepared %d packets, want %d", len(got), numPkts)
|
||||
}
|
||||
for i, pkt := range got {
|
||||
if !slices.Equal(pkt, bufs[i]) {
|
||||
t.Errorf("entry %d bytes differ from bufs[%d]", i, i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestOffloadsDisabledOnLoopback is the offloads-off smoke test, the mirror
|
||||
// of TestGSOEngagesOnLoopback: a real socket built with Offloads=false must
|
||||
// skip the GSO/GRO probes entirely (even on kernels that support them) and
|
||||
// still deliver a same-destination batch as plain per-packet datagrams.
|
||||
func TestOffloadsDisabledOnLoopback(t *testing.T) {
|
||||
rx, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
|
||||
if err != nil {
|
||||
t.Fatalf("listen rx: %v", err)
|
||||
}
|
||||
defer rx.Close()
|
||||
dst := rx.LocalAddr().(*net.UDPAddr).AddrPort()
|
||||
|
||||
udpSettings := Settings{
|
||||
Listen: netip.MustParseAddrPort("127.0.0.1:0"),
|
||||
Multi: false,
|
||||
Batch: 8, // batch > 1 would enable GRO if Offloads did not gate it
|
||||
Offloads: false,
|
||||
}
|
||||
uc, err := NewListener(testLogger(), udpSettings)
|
||||
if err != nil {
|
||||
t.Fatalf("NewListener: %v", err)
|
||||
}
|
||||
sc := uc.(*StdConn)
|
||||
defer sc.Close()
|
||||
|
||||
if sc.bw.gsoSupported {
|
||||
t.Error("gsoSupported true with offloads disabled: probe was not skipped")
|
||||
}
|
||||
if sc.bw.cmsg != nil {
|
||||
t.Error("cmsg slab allocated with offloads disabled")
|
||||
}
|
||||
if sc.groSupported {
|
||||
t.Error("groSupported true with offloads disabled: probe was not skipped")
|
||||
}
|
||||
|
||||
var entryCounts []int
|
||||
real := sc.bw.sendFn
|
||||
sc.bw.sendFn = func(start, n int) (int, error) {
|
||||
entryCounts = append(entryCounts, n)
|
||||
return real(start, n)
|
||||
}
|
||||
|
||||
const numPkts = 8
|
||||
const pktLen = 1200
|
||||
bufs := make([][]byte, numPkts)
|
||||
addrs := make([]netip.AddrPort, numPkts)
|
||||
for i := range bufs {
|
||||
bufs[i] = make([]byte, pktLen)
|
||||
for j := range bufs[i] {
|
||||
bufs[i][j] = byte(i)
|
||||
}
|
||||
addrs[i] = dst
|
||||
}
|
||||
|
||||
written, err := sc.WriteBatch(bufs, addrs)
|
||||
if err != nil {
|
||||
t.Fatalf("WriteBatch: %v", err)
|
||||
}
|
||||
if written != numPkts {
|
||||
t.Fatalf("written = %d, want %d", written, numPkts)
|
||||
}
|
||||
// One sendmmsg call with one entry per packet: a single-entry call here
|
||||
// means GSO engaged despite being disabled.
|
||||
if len(entryCounts) != 1 || entryCounts[0] != numPkts {
|
||||
t.Fatalf("sendmmsg entry counts = %v, want [%d]", entryCounts, numPkts)
|
||||
}
|
||||
|
||||
_ = rx.SetReadDeadline(time.Now().Add(5 * time.Second))
|
||||
got := make([]byte, pktLen+1)
|
||||
for i := range numPkts {
|
||||
n, _, err := rx.ReadFromUDP(got)
|
||||
if err != nil {
|
||||
t.Fatalf("rx read %d: %v", i, err)
|
||||
}
|
||||
if n != pktLen {
|
||||
t.Fatalf("rx read %d: len=%d want %d", i, n, pktLen)
|
||||
}
|
||||
for j := range n {
|
||||
if got[j] != byte(i) {
|
||||
t.Fatalf("rx read %d: byte %d = %#x, want %#x", i, j, got[j], byte(i))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestOffloadsDisabledRxDelivers exercises the receive path with GRO gated
|
||||
// off but batch reads still on: ListenOut must deliver plain datagrams via
|
||||
// the MTU-sized buffer layout (no cmsg slots).
|
||||
func TestOffloadsDisabledRxDelivers(t *testing.T) {
|
||||
udpSettings := Settings{
|
||||
Listen: netip.MustParseAddrPort("127.0.0.1:0"),
|
||||
Multi: false,
|
||||
Batch: 8,
|
||||
Offloads: false,
|
||||
}
|
||||
uc, err := NewListener(testLogger(), udpSettings)
|
||||
if err != nil {
|
||||
t.Fatalf("NewListener: %v", err)
|
||||
}
|
||||
sc := uc.(*StdConn)
|
||||
|
||||
addr, err := sc.LocalAddr()
|
||||
if err != nil {
|
||||
t.Fatalf("LocalAddr: %v", err)
|
||||
}
|
||||
|
||||
type rxPkt struct {
|
||||
from netip.AddrPort
|
||||
payload []byte
|
||||
}
|
||||
rxCh := make(chan rxPkt, 16)
|
||||
listenDone := make(chan struct{})
|
||||
go func() {
|
||||
defer close(listenDone)
|
||||
_ = sc.ListenOut(func(from netip.AddrPort, payload []byte) {
|
||||
// payload aliases the shared recv buffer row; copy before handing off.
|
||||
rxCh <- rxPkt{from, slices.Clone(payload)}
|
||||
}, func() {})
|
||||
}()
|
||||
|
||||
tx, err := net.DialUDP("udp4", nil, net.UDPAddrFromAddrPort(addr))
|
||||
if err != nil {
|
||||
t.Fatalf("dial tx: %v", err)
|
||||
}
|
||||
defer tx.Close()
|
||||
|
||||
want := [][]byte{
|
||||
[]byte("one"),
|
||||
make([]byte, 1200),
|
||||
make([]byte, 9000), // near-MTU datagram must fit the non-GRO buffer size
|
||||
}
|
||||
for i := range want[1] {
|
||||
want[1][i] = 0xAB
|
||||
}
|
||||
for i := range want[2] {
|
||||
want[2][i] = 0xCD
|
||||
}
|
||||
for i, p := range want {
|
||||
if _, err := tx.Write(p); err != nil {
|
||||
t.Fatalf("tx write %d: %v", i, err)
|
||||
}
|
||||
}
|
||||
|
||||
for i, p := range want {
|
||||
select {
|
||||
case got := <-rxCh:
|
||||
if !slices.Equal(got.payload, p) {
|
||||
t.Errorf("packet %d: payload differs (len=%d want %d)", i, len(got.payload), len(p))
|
||||
}
|
||||
if got.from.Port() != tx.LocalAddr().(*net.UDPAddr).AddrPort().Port() {
|
||||
t.Errorf("packet %d: from=%v, want sender port %d", i, got.from, tx.LocalAddr().(*net.UDPAddr).Port)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatalf("timed out waiting for packet %d", i)
|
||||
}
|
||||
}
|
||||
|
||||
if err := sc.Close(); err != nil {
|
||||
t.Fatalf("Close: %v", err)
|
||||
}
|
||||
select {
|
||||
case <-listenDone:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("ListenOut did not return after Close")
|
||||
}
|
||||
}
|
||||
+18
-12
@@ -5,10 +5,8 @@ package udp
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"runtime"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
@@ -17,16 +15,18 @@ import (
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
func testLogger() *slog.Logger {
|
||||
return slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: slog.LevelError}))
|
||||
}
|
||||
|
||||
// TestShutdownWakesAfterRx_Mechanism exercises the kernel quirk our teardown
|
||||
// relies on: once a socket has received a packet, shutdown(2) wakes a blocked
|
||||
// recvmmsg with n>=1/Len==0 (not n==0). recvmmsg must turn that into net.ErrClosed
|
||||
// once Close set closed, so a parked reader exits instead of spinning.
|
||||
func TestShutdownWakesAfterRx_Mechanism(t *testing.T) {
|
||||
c, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, true, 64)
|
||||
udpSettings := Settings{
|
||||
Listen: netip.MustParseAddrPort("127.0.0.1:0"),
|
||||
Multi: true,
|
||||
Batch: 64,
|
||||
Offloads: true,
|
||||
}
|
||||
c, err := NewListener(testLogger(), udpSettings)
|
||||
if err != nil {
|
||||
t.Fatalf("NewListener: %v", err)
|
||||
}
|
||||
@@ -35,7 +35,7 @@ func TestShutdownWakesAfterRx_Mechanism(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("LocalAddr: %v", err)
|
||||
}
|
||||
msgs, _, _ := sc.PrepareRawMessages(sc.batch)
|
||||
msgs, _, _, _ := prepareRawMessages(sc.batch, 0xffff, 16)
|
||||
|
||||
// Receive a real packet so the socket has carried data.
|
||||
send, err := net.Dial("udp", addr.String())
|
||||
@@ -109,8 +109,8 @@ func TestListenOutTeardown_TrafficPatterns(t *testing.T) {
|
||||
}},
|
||||
}
|
||||
|
||||
// batch 1 exercises the recvmsg path, batch 64 the recvmmsg path; both must
|
||||
// tear down cleanly.
|
||||
// batch 1 exercises single-message reads, batch 64 a full recvmmsg batch;
|
||||
// both must tear down cleanly.
|
||||
for _, batch := range []int{1, 64} {
|
||||
for _, tc := range cases {
|
||||
t.Run(fmt.Sprintf("batch%d/%s", batch, tc.name), func(t *testing.T) {
|
||||
@@ -121,7 +121,13 @@ func TestListenOutTeardown_TrafficPatterns(t *testing.T) {
|
||||
}
|
||||
|
||||
func runTeardownCase(t *testing.T, batch int, name string, traffic func(send net.Conn, stop <-chan struct{})) {
|
||||
c, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, true, batch)
|
||||
udpSettings := Settings{
|
||||
Listen: netip.MustParseAddrPort("127.0.0.1:0"),
|
||||
Multi: true,
|
||||
Batch: batch,
|
||||
Offloads: true,
|
||||
}
|
||||
c, err := NewListener(testLogger(), udpSettings)
|
||||
if err != nil {
|
||||
t.Fatalf("NewListener: %v", err)
|
||||
}
|
||||
@@ -136,7 +142,7 @@ func runTeardownCase(t *testing.T, batch int, name string, traffic func(send net
|
||||
go func() {
|
||||
loopDone <- sc.ListenOut(func(netip.AddrPort, []byte) {
|
||||
received.Add(1)
|
||||
})
|
||||
}, func() {})
|
||||
}()
|
||||
|
||||
send, err := net.Dial("udp", addr.String())
|
||||
|
||||
@@ -0,0 +1,387 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"strconv"
|
||||
"strings"
|
||||
"unsafe"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// batchWriter owns the sendmmsg(2)/UDP-GSO transmit path for a StdConn: the
|
||||
// scratch WriteBatch packs mmsghdr entries into, plus the GSO capability
|
||||
// state probed at socket creation. Each queue has its own StdConn and
|
||||
// batchWriter, so no locking is needed.
|
||||
//
|
||||
// Terminology, smallest to largest:
|
||||
//
|
||||
// packet one element of bufs: a single UDP datagram. The unit of the
|
||||
// returned written count.
|
||||
// run consecutive packets planRun groups into one entry: same
|
||||
// destination, equal sizes (a shorter packet only last), within
|
||||
// maxGSOBytes and maxGSOSegments. Without GSO a run is always one
|
||||
// packet. Runs are atomic: packed whole into one entry, or
|
||||
// skipped whole if the socket cannot address their destination,
|
||||
// leaving a hole (bufs indices covered by no entry).
|
||||
// entry one mmsghdr slot of the sendmmsg array; the kernel's unit of
|
||||
// success and failure. A multi-packet entry carries a UDP_SEGMENT
|
||||
// cmsg and is sent as one superpacket the kernel segments into
|
||||
// gso_size-byte datagrams. Entries never split.
|
||||
// chunk the entries packed for one sendmmsg call, at most MaxWriteBatch.
|
||||
// batch the caller's whole bufs/addrs pair, processed as one or more chunks.
|
||||
type batchWriter struct {
|
||||
fd int
|
||||
isV4 bool
|
||||
|
||||
// UDP GSO (sendmsg with UDP_SEGMENT cmsg) support, probed once at
|
||||
// socket creation and cleared by WriteBatch if the kernel later rejects
|
||||
// a GSO send (the setsockopt probe cannot see per-route limitations).
|
||||
// When true, WriteBatch coalesces runs into UDP_SEGMENT entries;
|
||||
// otherwise each packet is its own entry.
|
||||
gsoSupported bool
|
||||
maxGSOSegments int
|
||||
|
||||
// sendmmsg scratch, sized to MaxWriteBatch at construction; WriteBatch
|
||||
// chunks larger inputs.
|
||||
msgs []rawMessage
|
||||
iovs []iovec
|
||||
names [][]byte
|
||||
l *slog.Logger
|
||||
|
||||
// sendFn sends n prepared entries beginning at w.msgs[start]
|
||||
// This is a function pointer to facilitate testing.
|
||||
sendFn func(start, n int) (int, error)
|
||||
|
||||
// Per-entry cmsg scratch: one contiguous slab of
|
||||
// MaxWriteBatch * cmsgSpace bytes holding one UDP_SEGMENT cmsg per entry.
|
||||
cmsg []byte
|
||||
cmsgSpace int
|
||||
|
||||
// entryEnd[e] is the bufs index after the last packet packed into entry e.
|
||||
// entryEnd[e]-entryPkts[e] recovers the bufs index the entry's run started at,
|
||||
// used to rewind i for the GSO-disable replay.
|
||||
entryEnd []int
|
||||
|
||||
// entryPkts[e] is the number of packets packed into entry e.
|
||||
entryPkts []int
|
||||
}
|
||||
|
||||
func newBatchWriter(fd int, isV4 bool, l *slog.Logger, offloadsEnabled bool) *batchWriter {
|
||||
w := &batchWriter{fd: fd, isV4: isV4, l: l}
|
||||
w.sendFn = w.sendmmsg
|
||||
if offloadsEnabled {
|
||||
w.prepareGSO()
|
||||
}
|
||||
w.prepareWriteMessages(MaxWriteBatch, offloadsEnabled)
|
||||
return w
|
||||
}
|
||||
|
||||
// prepareWriteMessages allocates the per-entry mmsghdr/iovec/sockaddr/cmsg
|
||||
// scratch. Hdr.Iov/Iovlen/Control/Controllen are wired per call, since an
|
||||
// entry spans a variable number of iovecs and may or may not carry a cmsg.
|
||||
//
|
||||
// Each entry's cmsg slot holds one UDP_SEGMENT (gso_size, uint16) header,
|
||||
// pre-filled here; only its payload is rewritten per call.
|
||||
// Hdr.Control/Controllen select whether it applies (none / segment).
|
||||
func (w *batchWriter) prepareWriteMessages(n int, offloadsEnabled bool) {
|
||||
w.msgs = make([]rawMessage, n)
|
||||
w.iovs = make([]iovec, n)
|
||||
w.names = make([][]byte, n)
|
||||
w.entryEnd = make([]int, n)
|
||||
w.entryPkts = make([]int, n)
|
||||
|
||||
w.cmsgSpace = unix.CmsgSpace(2)
|
||||
|
||||
for i := range w.msgs {
|
||||
w.names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||
w.msgs[i].Hdr.Name = &w.names[i][0]
|
||||
}
|
||||
|
||||
if !offloadsEnabled || !w.gsoSupported {
|
||||
return //avoid allocating cmsg space if we will never use it
|
||||
}
|
||||
|
||||
w.cmsg = make([]byte, n*w.cmsgSpace)
|
||||
|
||||
for k := 0; k < n; k++ {
|
||||
base := k * w.cmsgSpace
|
||||
seg := (*unix.Cmsghdr)(unsafe.Pointer(&w.cmsg[base]))
|
||||
seg.Level = unix.SOL_UDP
|
||||
seg.Type = unix.UDP_SEGMENT
|
||||
setCmsgLen(seg, unix.CmsgLen(2))
|
||||
}
|
||||
}
|
||||
|
||||
// maxGSOBytes bounds the total payload of one UDP_SEGMENT send. The kernel
|
||||
// builds a single skb, which must fit the 16-bit UDP length field and
|
||||
// sk_gso_max_size (65536 on most devices); 65000 leaves headroom for headers.
|
||||
const maxGSOBytes = 65000
|
||||
|
||||
// prepareGSO probes UDP_SEGMENT support and sets w.gsoSupported on success.
|
||||
// Best-effort; failure leaves it false.
|
||||
func (w *batchWriter) prepareGSO() {
|
||||
w.maxGSOSegments = 63 // pre-6.9 cap; see gsoMaxSegments
|
||||
|
||||
if err := unix.SetsockoptInt(w.fd, unix.IPPROTO_UDP, unix.UDP_SEGMENT, 0); err != nil {
|
||||
w.l.Info("udp: GSO disabled", "reason", "rawconn control failed", "error", err)
|
||||
recordCapability("udp.gso.enabled", false)
|
||||
return
|
||||
}
|
||||
|
||||
var un unix.Utsname
|
||||
if err := unix.Uname(&un); err != nil {
|
||||
w.l.Warn("udp: kernel version probe failed, capping GSO at 63 segments", "error", err)
|
||||
} else {
|
||||
w.maxGSOSegments = gsoMaxSegments(string(un.Release[:]))
|
||||
}
|
||||
|
||||
w.gsoSupported = true
|
||||
w.l.Info("udp: GSO enabled", "maxGSOSegments", w.maxGSOSegments)
|
||||
recordCapability("udp.gso.enabled", true)
|
||||
}
|
||||
|
||||
// gsoMaxSegments returns the most segments one UDP_SEGMENT send may carry:
|
||||
// the kernel cap (UDP_MAX_SEGMENTS: 64 before 6.9, 128 after) minus one,
|
||||
// because the kernel counts the 8-byte UDP header against the gso_size * UDP_MAX_SEGMENTS budget.
|
||||
func gsoMaxSegments(release string) int {
|
||||
major, minor := parseRelease(release)
|
||||
if major > 6 || (major == 6 && minor >= 9) {
|
||||
return 127
|
||||
}
|
||||
return 63
|
||||
}
|
||||
|
||||
func parseRelease(r string) (major, minor int) {
|
||||
// strip anything after the second dot or any non-digit
|
||||
parts := strings.SplitN(r, ".", 3)
|
||||
if len(parts) < 2 {
|
||||
return 0, 0
|
||||
}
|
||||
major, _ = strconv.Atoi(parts[0])
|
||||
// minor may have trailing junk like "15-generic"
|
||||
mp := parts[1]
|
||||
for i, c := range mp {
|
||||
if c < '0' || c > '9' {
|
||||
mp = mp[:i]
|
||||
break
|
||||
}
|
||||
}
|
||||
minor, _ = strconv.Atoi(mp)
|
||||
return
|
||||
}
|
||||
|
||||
// WriteBatch sends bufs via sendmmsg(2), coalescing runs into UDP_SEGMENT
|
||||
// entries, so one syscall can mix GSO superpackets and plain datagrams.
|
||||
// Without GSO support every packet is its own entry.
|
||||
// Callers shall deliver same-destination packets contiguously and in counter order
|
||||
//
|
||||
// Batches larger than the scratch take one sendmmsg per chunk.
|
||||
// A partial success resumes the same prepared entries at the first unsent entry.
|
||||
// A zero-sent error means the kernel rejected the first remaining entry:
|
||||
// its packets are dropped and the rest of the chunk resumes in place.
|
||||
//
|
||||
// Returns the number of packets sent. An error means the call itself failed.
|
||||
// A short count means some destinations were undeliverable.
|
||||
func (w *batchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
|
||||
if len(bufs) != len(addrs) {
|
||||
return 0, fmt.Errorf("WriteBatch: len(bufs)=%d != len(addrs)=%d", len(bufs), len(addrs))
|
||||
}
|
||||
|
||||
// A destination the kernel rejects results in us dropping that entry (one packet, or one same-destination GSO run).
|
||||
// We count what actually made it out rather than returning an error.
|
||||
written := 0
|
||||
|
||||
i := 0
|
||||
for i < len(bufs) {
|
||||
entry := 0
|
||||
iovIdx := 0
|
||||
for entry < len(w.msgs) && i < len(bufs) {
|
||||
iovBudget := len(w.iovs) - iovIdx
|
||||
if iovBudget < 1 {
|
||||
break
|
||||
}
|
||||
runLen, segSize := w.planRun(bufs, addrs, i, iovBudget)
|
||||
if runLen == 0 {
|
||||
break
|
||||
}
|
||||
|
||||
for k := 0; k < runLen; k++ {
|
||||
b := bufs[i+k]
|
||||
if len(b) == 0 {
|
||||
w.iovs[iovIdx+k].Base = nil
|
||||
setIovLen(&w.iovs[iovIdx+k], 0)
|
||||
} else {
|
||||
w.iovs[iovIdx+k].Base = &b[0]
|
||||
setIovLen(&w.iovs[iovIdx+k], len(b))
|
||||
}
|
||||
}
|
||||
|
||||
nlen, err := writeSockaddr(w.names[entry], addrs[i], w.isV4)
|
||||
if err != nil {
|
||||
// The destination's address family does not match the socket
|
||||
// (e.g. an IPv6 remote on a v4-bound socket). The packets are
|
||||
// undeliverable and no entry is committed yet: skip the run.
|
||||
if w.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
w.l.Debug("skipping unroutable batch destination", "udpAddr", addrs[i], "packets", runLen, "error", err)
|
||||
}
|
||||
i += runLen
|
||||
continue
|
||||
}
|
||||
|
||||
hdr := &w.msgs[entry].Hdr
|
||||
hdr.Iov = &w.iovs[iovIdx]
|
||||
setMsgIovlen(hdr, runLen)
|
||||
hdr.Namelen = uint32(nlen)
|
||||
|
||||
w.writeEntryCmsg(entry, runLen, segSize)
|
||||
|
||||
i += runLen
|
||||
iovIdx += runLen
|
||||
w.entryEnd[entry] = i
|
||||
w.entryPkts[entry] = runLen
|
||||
entry++
|
||||
}
|
||||
|
||||
if entry == 0 {
|
||||
// Every remaining packet was skipped; i reached len(bufs).
|
||||
break
|
||||
}
|
||||
|
||||
// Drain the packed entries without repacking: everything the packing
|
||||
// loop wired (iovecs, names, cmsgs) stays intact until the next chunk
|
||||
// overwrites it, so a partial success resumes the same sendmmsg array
|
||||
// at the first unsent entry, and a rejected entry is skipped in place.
|
||||
// Only the GSO-disable path replans, since its entries change shape.
|
||||
done := 0
|
||||
for done < entry {
|
||||
sent, serr := w.sendFn(done, entry-done)
|
||||
if sent > 0 {
|
||||
// Count packets per entry; the bufs index span would
|
||||
// overcount across holes left by skipped runs.
|
||||
for e := done; e < done+sent; e++ {
|
||||
written += w.entryPkts[e]
|
||||
}
|
||||
done += sent
|
||||
continue
|
||||
}
|
||||
if serr == nil {
|
||||
return written, fmt.Errorf("sendmmsg made no progress")
|
||||
}
|
||||
// sent<=0 means the first remaining entry itself failed.
|
||||
// EIO on a superpacket means the route cannot carry a GSO send even though the setsockopt probe passed:
|
||||
// udp_send_skb() returns EIO when:
|
||||
// * the egress device lacks TX checksum offload (kernels through 6.10)
|
||||
// * or when an xfrm policy covers the route.
|
||||
// Persistent, so disable GSO and replay from the failed run as one-packet entries.
|
||||
if w.gsoSupported && w.entryPkts[done] >= 2 && errors.Is(serr, unix.EIO) {
|
||||
w.gsoSupported = false
|
||||
w.l.Warn("udp: kernel rejected GSO send, disabling GSO", "error", serr)
|
||||
recordCapability("udp.gso.enabled", false)
|
||||
i = w.entryEnd[done] - w.entryPkts[done]
|
||||
break
|
||||
}
|
||||
// Any other zero-sent error is a per-entry failure.
|
||||
// Transient errnos (EINTR, ENOBUFS) were already retried inside sendFn.
|
||||
// These packets are doomed. Log them and move on.
|
||||
if w.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
w.l.Debug("sendmmsg rejected entry",
|
||||
"error", serr,
|
||||
"udpAddr", addrs[w.entryEnd[done]-w.entryPkts[done]],
|
||||
"packets", w.entryPkts[done],
|
||||
"gso", w.gsoSupported,
|
||||
)
|
||||
}
|
||||
done++
|
||||
}
|
||||
}
|
||||
return written, nil
|
||||
}
|
||||
|
||||
// planRun returns the length of the run starting at start and its segment
|
||||
// size (len(bufs[start])). A run of length 1 carries no UDP_SEGMENT cmsg
|
||||
// and is sent as a plain datagram; without GSO support planRun always returns 1.
|
||||
func (w *batchWriter) planRun(bufs [][]byte, addrs []netip.AddrPort, start, iovBudget int) (int, int) {
|
||||
if start >= len(bufs) || iovBudget < 1 {
|
||||
return 0, 0
|
||||
}
|
||||
segSize := len(bufs[start])
|
||||
if !w.gsoSupported || segSize == 0 || segSize > maxGSOBytes {
|
||||
return 1, segSize
|
||||
}
|
||||
dst := addrs[start]
|
||||
maxLen := w.maxGSOSegments
|
||||
if iovBudget < maxLen {
|
||||
maxLen = iovBudget
|
||||
}
|
||||
runLen := 1
|
||||
total := segSize
|
||||
for runLen < maxLen && start+runLen < len(bufs) {
|
||||
nextLen := len(bufs[start+runLen])
|
||||
if nextLen == 0 || nextLen > segSize {
|
||||
break
|
||||
}
|
||||
if addrs[start+runLen] != dst {
|
||||
break
|
||||
}
|
||||
if total+nextLen > maxGSOBytes {
|
||||
break
|
||||
}
|
||||
total += nextLen
|
||||
runLen++
|
||||
if nextLen < segSize {
|
||||
// A short packet must be the last in the run.
|
||||
break
|
||||
}
|
||||
}
|
||||
return runLen, segSize
|
||||
}
|
||||
|
||||
// writeEntryCmsg writes one entry's UDP_SEGMENT payload when runLen >= 2 and
|
||||
// points Hdr.Control at it; a single-packet entry carries no cmsg.
|
||||
func (w *batchWriter) writeEntryCmsg(entry, runLen, segSize int) {
|
||||
hdr := &w.msgs[entry].Hdr
|
||||
base := entry * w.cmsgSpace
|
||||
|
||||
if runLen >= 2 {
|
||||
dataOff := base + unix.CmsgLen(0)
|
||||
binary.NativeEndian.PutUint16(w.cmsg[dataOff:dataOff+2], uint16(segSize))
|
||||
hdr.Control = &w.cmsg[base]
|
||||
setMsgControllen(hdr, w.cmsgSpace)
|
||||
} else {
|
||||
hdr.Control = nil
|
||||
setMsgControllen(hdr, 0)
|
||||
}
|
||||
}
|
||||
|
||||
// sendmmsg issues sendmmsg(2) against n entries of w.msgs starting at start.
|
||||
//
|
||||
// EINTR is automatically retried and will never be returned.
|
||||
// ENOBUFS is retried enobufsRetries times, and should be treated like any other error
|
||||
func (w *batchWriter) sendmmsg(start, n int) (int, error) {
|
||||
const enobufsRetries = 3
|
||||
for enobufs := 0; ; {
|
||||
r1, _, errno := unix.Syscall6(unix.SYS_SENDMMSG, uintptr(w.fd),
|
||||
uintptr(unsafe.Pointer(&w.msgs[start])), uintptr(n),
|
||||
0, 0, 0,
|
||||
)
|
||||
switch {
|
||||
case errno == unix.EINTR: //similar to stdlib's ignoringEINTRIO
|
||||
continue
|
||||
case errno == unix.ENOBUFS && enobufs < enobufsRetries:
|
||||
enobufs++ //worth a retry or three
|
||||
continue
|
||||
case errno != 0:
|
||||
return int(r1), &net.OpError{Op: "sendmmsg", Err: errno}
|
||||
}
|
||||
return int(r1), nil
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// TestWriteBatchNoAllocs verifies the sendmmsg/UDP-GSO transmit path performs
|
||||
// no per-packet heap allocations on the happy path: all mmsghdr/iovec/cmsg
|
||||
// scratch is preallocated in newBatchWriter and WriteBatch may only rewrite
|
||||
// it. The batch deliberately mixes a GSO-eligible run, a short tail segment,
|
||||
// destination changes, so the planner, sockaddr,
|
||||
// and cmsg paths are all exercised.
|
||||
func TestWriteBatchNoAllocs(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
addr string
|
||||
}{
|
||||
{"v4", "127.0.0.1"},
|
||||
{"v6", "::1"},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
ip := netip.MustParseAddr(tc.addr)
|
||||
newConn := func() Conn {
|
||||
udpSettings := Settings{
|
||||
Listen: netip.AddrPortFrom(ip, 0),
|
||||
Multi: false,
|
||||
Batch: 8,
|
||||
Offloads: true,
|
||||
}
|
||||
c, err := NewListener(testLogger(), udpSettings)
|
||||
if err != nil {
|
||||
t.Fatalf("NewListener: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = c.Close() })
|
||||
return c
|
||||
}
|
||||
tx := newConn()
|
||||
rxA := newConn()
|
||||
rxB := newConn()
|
||||
if sc, ok := tx.(*StdConn); ok {
|
||||
// Records which planner path the measurement covered; GSO
|
||||
// support depends on the running kernel.
|
||||
t.Logf("gsoSupported=%v maxGSOSegments=%d", sc.bw.gsoSupported, sc.bw.maxGSOSegments)
|
||||
}
|
||||
dstA, err := rxA.LocalAddr()
|
||||
if err != nil {
|
||||
t.Fatalf("LocalAddr: %v", err)
|
||||
}
|
||||
dstB, err := rxB.LocalAddr()
|
||||
if err != nil {
|
||||
t.Fatalf("LocalAddr: %v", err)
|
||||
}
|
||||
|
||||
payload := make([]byte, 1200)
|
||||
short := make([]byte, 900)
|
||||
|
||||
var bufs [][]byte
|
||||
var addrs []netip.AddrPort
|
||||
add := func(b []byte, dst netip.AddrPort) {
|
||||
bufs = append(bufs, b)
|
||||
addrs = append(addrs, dst)
|
||||
}
|
||||
// GSO-eligible run with a short tail.
|
||||
for k := 0; k < 8; k++ {
|
||||
add(payload, dstA)
|
||||
}
|
||||
add(short, dstA)
|
||||
add(payload, dstA)
|
||||
// Alternating destinations defeat coalescing entirely.
|
||||
for k := 0; k < 4; k++ {
|
||||
dst := dstA
|
||||
if k%2 == 0 {
|
||||
dst = dstB
|
||||
}
|
||||
add(payload, dst)
|
||||
}
|
||||
|
||||
var werr error
|
||||
// Warm-up outside the measured runs.
|
||||
if _, err := tx.WriteBatch(bufs, addrs); err != nil {
|
||||
t.Fatalf("WriteBatch warm-up: %v", err)
|
||||
}
|
||||
allocs := testing.AllocsPerRun(100, func() {
|
||||
if _, err := tx.WriteBatch(bufs, addrs); err != nil {
|
||||
werr = err
|
||||
}
|
||||
})
|
||||
if werr != nil {
|
||||
t.Fatalf("WriteBatch: %v", werr)
|
||||
}
|
||||
if allocs != 0 {
|
||||
t.Fatalf("WriteBatch allocated %.1f times per call, want 0", allocs)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
+2
-3
@@ -9,14 +9,13 @@ import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"syscall"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) {
|
||||
return NewGenericListener(l, ip, port, multi, batch)
|
||||
func NewListener(l *slog.Logger, s Settings) (Conn, error) {
|
||||
return NewGenericListener(l, s)
|
||||
}
|
||||
|
||||
func NewListenConfig(multi bool) net.ListenConfig {
|
||||
|
||||
+16
-2
@@ -140,7 +140,7 @@ func (u *RIOConn) bind(l *slog.Logger, sa windows.Sockaddr) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (u *RIOConn) ListenOut(r EncReader) error {
|
||||
func (u *RIOConn) ListenOut(r EncReader, flush func()) error {
|
||||
buffer := make([]byte, MTU)
|
||||
|
||||
var lastRecvErr time.Time
|
||||
@@ -161,7 +161,8 @@ func (u *RIOConn) ListenOut(r EncReader) error {
|
||||
continue
|
||||
}
|
||||
|
||||
r(netip.AddrPortFrom(netip.AddrFrom16(rua.Addr).Unmap(), (rua.Port>>8)|((rua.Port&0xff)<<8)), buffer[:n])
|
||||
r(netip.AddrPortFrom(netip.AddrFrom16(rua.Addr).Unmap(), (rua.Port>>8)|((rua.Port&0xff)<<8)), buffer[:n:n])
|
||||
flush()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -316,6 +317,19 @@ func (u *RIOConn) WriteTo(buf []byte, ip netip.AddrPort) error {
|
||||
return winrio.SendEx(u.rq, dataBuffer, 1, nil, addressBuffer, nil, nil, 0, 0)
|
||||
}
|
||||
|
||||
func (u *RIOConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
|
||||
// An un-sendable destination costs its own packet, never the ones behind it in the batch.
|
||||
written := 0
|
||||
for i, b := range bufs {
|
||||
if err := u.WriteTo(b, addrs[i]); err == nil {
|
||||
written++
|
||||
} else {
|
||||
u.l.Debug("failed to write packet in batch", "udpAddr", addrs[i], "error", err)
|
||||
}
|
||||
}
|
||||
return written, nil
|
||||
}
|
||||
|
||||
func (u *RIOConn) LocalAddr() (netip.AddrPort, error) {
|
||||
sa, err := windows.Getsockname(u.sock)
|
||||
if err != nil {
|
||||
|
||||
+18
-4
@@ -84,14 +84,14 @@ type TesterConn struct {
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
func NewListener(l *slog.Logger, ip netip.Addr, port int, _ bool, _ int) (Conn, error) {
|
||||
func NewListener(l *slog.Logger, s Settings) (Conn, error) {
|
||||
c := &TesterConn{
|
||||
RxPackets: make(chan *Packet, 10),
|
||||
TxPackets: make(chan *Packet, 10),
|
||||
done: make(chan struct{}),
|
||||
l: l,
|
||||
}
|
||||
c.SetAddr(netip.AddrPortFrom(ip, uint16(port)))
|
||||
c.SetAddr(s.Listen)
|
||||
return c, nil
|
||||
}
|
||||
|
||||
@@ -171,14 +171,28 @@ func (u *TesterConn) WriteTo(b []byte, addr netip.AddrPort) error {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
func (u *TesterConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, error) {
|
||||
written := 0
|
||||
for i, b := range bufs {
|
||||
if err := u.WriteTo(b, addrs[i]); err == nil {
|
||||
written++
|
||||
} else {
|
||||
u.l.Debug("failed to write packet in batch", "udpAddr", addrs[i], "error", err)
|
||||
}
|
||||
}
|
||||
return written, nil
|
||||
}
|
||||
|
||||
func (u *TesterConn) ListenOut(r EncReader) error {
|
||||
func (u *TesterConn) ListenOut(r EncReader, flush func()) error {
|
||||
for {
|
||||
select {
|
||||
case <-u.done:
|
||||
return os.ErrClosed
|
||||
case p := <-u.RxPackets:
|
||||
r(p.From, p.Data)
|
||||
r(p.From, p.Data[:len(p.Data):len(p.Data)])
|
||||
// The batcher borrows plaintext decrypted in place inside p.Data
|
||||
// until Flush, so the packet must stay alive across flush()
|
||||
flush()
|
||||
p.Release()
|
||||
}
|
||||
}
|
||||
|
||||
+4
-5
@@ -7,12 +7,11 @@ import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) {
|
||||
if multi {
|
||||
func NewListener(l *slog.Logger, s Settings) (Conn, error) {
|
||||
if s.Multi {
|
||||
//NOTE: Technically we can support it with RIO but it wouldn't be at the socket level
|
||||
// The udp stack would need to be reworked to hide away the implementation differences between
|
||||
// Windows and Linux
|
||||
@@ -20,12 +19,12 @@ func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int)
|
||||
}
|
||||
|
||||
var conn Conn
|
||||
rc, err := NewRIOListener(l, ip, port)
|
||||
rc, err := NewRIOListener(l, s.Listen.Addr(), int(s.Listen.Port()))
|
||||
if err == nil {
|
||||
conn = rc
|
||||
} else {
|
||||
l.Error("Falling back to standard udp sockets", "error", err)
|
||||
conn, err = NewGenericListener(l, ip, port, multi, batch)
|
||||
conn, err = NewGenericListener(l, s)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user