Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
131 changes: 105 additions & 26 deletions internal/socket/coverage_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import (
"errors"
"net"
"net/netip"
"slices"
"strings"
"testing"
"time"
Expand Down Expand Up @@ -1114,25 +1115,36 @@ func TestRelayTransportServeForwardsRecv(t *testing.T) {
Datagrams: relayproto.Datagrams{SegmentSize: 2, Contents: []byte("aabbcc")},
}

want := []string{"aa", "bb", "cc"}
deadline := time.After(5 * time.Second)
for i, w := range want {
select {
case b := <-recvCh:
if string(b.data) != w {
t.Errorf("segment %d = %q, want %q", i, b.data, w)
}
// Each batch must be tagged with the relay Addr it arrived on.
gu, ge, ok := b.info.Remote.Relay()
if !ok {
t.Errorf("segment %d Remote kind = %v, want relay", i, b.info.Remote.Kind())
} else if !gu.Equal(url) || !ge.Equal(src.Public().EndpointID()) {
t.Errorf("segment %d Remote = (%s, %s), want (%s, %s)", i, gu, ge, url, src.Public().EndpointID())
}
case <-deadline:
t.Fatalf("timed out waiting for segment %d", i)
select {
case b := <-recvCh:
if got := batchSegments(b); !slices.Equal(got, []string{"aa", "bb", "cc"}) {
t.Errorf("segments = %q", got)
}
// The batch must be tagged with the relay Addr it arrived on.
gu, ge, ok := b.info.Remote.Relay()
if !ok {
t.Errorf("Remote kind = %v, want relay", b.info.Remote.Kind())
} else if !gu.Equal(url) || !ge.Equal(src.Public().EndpointID()) {
t.Errorf("Remote = (%s, %s), want (%s, %s)", gu, ge, url, src.Public().EndpointID())
}
case <-time.After(5 * time.Second):
t.Fatal("timed out waiting for batch")
}
}

// batchSegments splits b the way MagicConn.ReadFrom does.
func batchSegments(b recvBatch) []string {
var out []string
d := b.data
for len(d) > 0 {
n := len(d)
if b.stride > 0 && n > b.stride {
n = b.stride
}
out = append(out, string(d[:n]))
d = d[n:]
}
return out
}

// TestRelayTransportDeliverSegments unit-tests deliver directly: an 8-byte
Expand All @@ -1152,16 +1164,13 @@ func TestRelayTransportDeliverSegments(t *testing.T) {
}

rt.deliver(context.Background(), dm)
wantLens := []int{3, 3, 2}
for i, wl := range wantLens {
select {
case b := <-recvCh:
if len(b.data) != wl {
t.Errorf("segment %d len = %d, want %d", i, len(b.data), wl)
}
default:
t.Fatalf("missing segment %d", i)
select {
case b := <-recvCh:
if got := batchSegments(b); !slices.Equal(got, []string{"aaa", "bbb", "cc"}) {
t.Errorf("segments = %q", got)
}
default:
t.Fatal("missing batch")
}

// A cancelled context makes deliver return before enqueuing into a full
Expand Down Expand Up @@ -1263,3 +1272,73 @@ func TestMagicConnRelayAccessor(t *testing.T) {

m.Close()
}

// TestReadFromDropsEmptyBatch checks an empty batch is released rather than
// left in the cursor, where the next batch would overwrite it unreleased.
func TestReadFromDropsEmptyBatch(t *testing.T) {
udp, err := net.ListenUDP("udp", net.UDPAddrFromAddrPort(netip.AddrPortFrom(netip.IPv6Loopback(), 0)))
if err != nil {
t.Fatal(err)
}
defer udp.Close()
m := NewMagicConnWithTransports(NewSocket(), udp, nil)
released := 0
src := netip.MustParseAddrPort("192.0.2.1:7")
m.recvCh <- recvBatch{ip: src, releaseFn: func() { released++ }}
m.recvCh <- recvBatch{data: []byte("hi"), ip: src, releaseFn: func() { released++ }}
if err := m.SetReadDeadline(time.Now().Add(5 * time.Second)); err != nil {
t.Fatal(err)
}
buf := make([]byte, 16)
n, _, err := m.ReadFrom(buf)
if err != nil {
t.Fatal(err)
}
if got := string(buf[:n]); got != "hi" {
t.Errorf("read %q, want %q", got, "hi")
}
if released != 2 {
t.Errorf("released = %d, want 2: the empty batch was not released", released)
}
}

// TestReadFromSplitsBatch checks ReadFrom hands out one segment per call from a
// strided recvBatch and releases the batch after the last one.
func TestReadFromSplitsBatch(t *testing.T) {
udp, err := net.ListenUDP("udp", net.UDPAddrFromAddrPort(netip.AddrPortFrom(netip.IPv6Loopback(), 0)))
if err != nil {
t.Fatal(err)
}
defer udp.Close()
m := NewMagicConnWithTransports(NewSocket(), udp, nil)
released := 0
src := netip.MustParseAddrPort("192.0.2.1:7")
m.recvCh <- recvBatch{data: []byte("aaabbbcc"), stride: 3, ip: src, releaseFn: func() { released++ }}
// Without a deadline a ReadFrom that handed out too much at once would
// block here instead of reporting the segment it got wrong.
if err := m.SetReadDeadline(time.Now().Add(5 * time.Second)); err != nil {
t.Fatal(err)
}
buf := make([]byte, 16)
for i, want := range []string{"aaa", "bbb", "cc"} {
n, addr, err := m.ReadFrom(buf)
if err != nil {
t.Fatalf("segment %d: %v", i, err)
}
if string(buf[:n]) != want {
t.Errorf("segment %d = %q, want %q", i, buf[:n], want)
}
if addr.(*net.UDPAddr).AddrPort() != src {
t.Errorf("segment %d addr = %v", i, addr)
}
if i < 2 && released != 0 {
t.Errorf("released before last segment")
}
}
if released != 1 {
t.Errorf("released = %d, want 1", released)
}
if got := m.Metrics().RecvDatagrams; got != 3 {
t.Errorf("RecvDatagrams = %d, want 3", got)
}
}
61 changes: 60 additions & 1 deletion internal/socket/ip.go
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,10 @@ type IpTransport struct {
conn *net.UDPConn
recvCh chan<- recvBatch

// gro is set when the kernel coalesces received datagrams for this
// socket; the receive loop then reads runs rather than datagrams.
gro bool

// pktinfo is set when the socket is bound to the unspecified address and
// the platform delivers the destination address of received datagrams.
// A wildcard socket has no fixed source address: the kernel picks one per
Expand Down Expand Up @@ -67,11 +71,17 @@ type localEntry struct {
// from the address the kernel picks, as before the table.
const maxLocalAddrs = 4096

// groBufSize bounds one UDP_GRO read: the kernel coalesces at most a 64 KiB
// run of datagrams into a single recvmsg.
const groBufSize = 65535

var groRecvPool = sync.Pool{New: func() any { b := make([]byte, groBufSize); return &b }}

// NewIpTransport returns an IpTransport over conn that delivers received
// datagrams to recvCh. The transport does not take ownership of conn; the caller
// closes it.
func NewIpTransport(conn *net.UDPConn, recvCh chan<- recvBatch) *IpTransport {
t := &IpTransport{conn: conn, recvCh: recvCh}
t := &IpTransport{conn: conn, recvCh: recvCh, gro: enableGRO(conn)}
if la, ok := conn.LocalAddr().(*net.UDPAddr); ok && la.IP.IsUnspecified() {
t.pktinfo = enablePacketInfo(conn)
if t.pktinfo {
Expand Down Expand Up @@ -191,6 +201,10 @@ func (t *IpTransport) LocalAddr() net.Addr { return t.conn.LocalAddr() }
// match iroh/src/socket/transports/ip.rs:221 to_canonical). Empty datagrams and
// transient errors are skipped; a closed socket ends the loop cleanly.
func (t *IpTransport) Serve(ctx context.Context) {
if t.gro {
t.serveGRO(ctx)
return
}
var oob []byte
if t.pktinfo {
oob = make([]byte, maxControlSize)
Expand Down Expand Up @@ -238,6 +252,51 @@ func (t *IpTransport) Serve(ctx context.Context) {
}
}

// serveGRO is the receive loop for a socket with UDP_GRO enabled: one read can
// return a run of equally sized datagrams from the same peer, so it queues the
// whole read as one batch strided by the segment size the kernel reports and
// hands the buffer back to the pool once ReadFrom has copied out the last
// segment.
func (t *IpTransport) serveGRO(ctx context.Context) {
// The read carries the segment size and, on a wildcard socket, the
// packet-info message too: room for both, not just the one.
var oob [groOOBSize + maxControlSize]byte
for {
if ctx.Err() != nil {
return
}
bp := groRecvPool.Get().(*[]byte)
buf := *bp
n, oobn, _, ap, err := t.conn.ReadMsgUDPAddrPort(buf, oob[:])
if err != nil {
groRecvPool.Put(bp)
if errors.Is(err, net.ErrClosed) || ctx.Err() != nil {
return
}
continue
}
if n == 0 {
groRecvPool.Put(bp)
continue
}
seg := groSegmentSize(oob[:oobn])
if seg <= 0 || seg > n {
seg = n
}
recordUDPReceive((n+seg-1)/seg, seg < n)
src := canonicalAddrPort(ap)
if t.pktinfo {
// A coalesced run is one peer's, so one arrival address
// covers every datagram in it.
t.recordLocal(src, oob[:oobn])
}
b := recvBatch{data: buf[:n], stride: seg, ip: src, groBuf: bp}
if !t.enqueue(ctx, b) {
return
}
}
}

func (t *IpTransport) enqueue(ctx context.Context, b recvBatch) bool {
select {
case t.recvCh <- b:
Expand Down
43 changes: 43 additions & 0 deletions internal/socket/ip_gro_linux.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,43 @@
//go:build linux

package socket

import (
"encoding/binary"
"net"

"golang.org/x/sys/unix"
)

// enableGRO turns on UDP_GRO so one recvmsg can return many coalesced
// datagrams from the same peer.
func enableGRO(conn *net.UDPConn) bool {
rc, err := conn.SyscallConn()
if err != nil {
return false
}
var serr error
if err := rc.Control(func(fd uintptr) {
serr = unix.SetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_GRO, 1)
}); err != nil {
return false
}
return serr == nil
}

const groOOBSize = 64

// groSegmentSize returns the UDP_GRO segment size in oob, or 0.
func groSegmentSize(oob []byte) int {
for len(oob) > 0 {
hdr, data, rest, err := unix.ParseOneSocketControlMessage(oob)
if err != nil {
return 0
}
if hdr.Level == unix.IPPROTO_UDP && hdr.Type == unix.UDP_GRO && len(data) >= 4 {
return int(int32(binary.NativeEndian.Uint32(data)))
}
oob = rest
}
return 0
}
Loading