diff --git a/internal/socket/coverage_test.go b/internal/socket/coverage_test.go index 67bbf888..63fb7dc6 100644 --- a/internal/socket/coverage_test.go +++ b/internal/socket/coverage_test.go @@ -5,6 +5,7 @@ import ( "errors" "net" "net/netip" + "slices" "strings" "testing" "time" @@ -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 @@ -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 @@ -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) + } +} diff --git a/internal/socket/ip.go b/internal/socket/ip.go index 1f8e5acd..93f3dbb6 100644 --- a/internal/socket/ip.go +++ b/internal/socket/ip.go @@ -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 @@ -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 { @@ -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) @@ -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: diff --git a/internal/socket/ip_gro_linux.go b/internal/socket/ip_gro_linux.go new file mode 100644 index 00000000..d6010c1c --- /dev/null +++ b/internal/socket/ip_gro_linux.go @@ -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 +} diff --git a/internal/socket/ip_gro_linux_test.go b/internal/socket/ip_gro_linux_test.go new file mode 100644 index 00000000..4c222f71 --- /dev/null +++ b/internal/socket/ip_gro_linux_test.go @@ -0,0 +1,190 @@ +//go:build linux + +package socket + +import ( + "bytes" + "context" + "net" + "net/netip" + "syscall" + "testing" + "time" + "unsafe" + + "golang.org/x/sys/unix" +) + +// TestIpTransportGROCoalesces exercises the UDP_GRO receive path against a +// UDP_SEGMENT sender on loopback: the kernel must hand the run of equally sized +// datagrams to one read, and Serve must queue it as a single strided recvBatch +// that ReadFrom can split back into the original datagrams. +func TestIpTransportGROCoalesces(t *testing.T) { + udp, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer udp.Close() + recvCh := make(chan recvBatch, 16) + tr := NewIpTransport(udp, recvCh) + if !tr.gro { + t.Skip("UDP_GRO not available on this kernel") + } + + sender, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer sender.Close() + + const ( + segSize = 1200 + segs = 8 + ) + payload := make([]byte, segSize*segs) + for i := range payload { + payload[i] = byte('a' + i/segSize) + } + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go tr.Serve(ctx) + + dst := udp.LocalAddr().(*net.UDPAddr) + if _, _, err := sender.WriteMsgUDP(payload, udpSegmentMessage(segSize), dst); err != nil { + t.Fatalf("segmented write: %v", err) + } + + var got []byte + coalesced := false + deadline := time.After(10 * time.Second) + for len(got) < len(payload) { + select { + case b := <-recvCh: + segs := batchSegments(b) + if len(segs) > 1 { + coalesced = true + } + for i, seg := range segs { + if len(seg) != segSize { + t.Errorf("segment %d is %d bytes, want %d", i, len(seg), segSize) + } + got = append(got, seg...) + } + b.release() + case <-deadline: + t.Fatalf("got %d of %d bytes", len(got), len(payload)) + } + } + if !bytes.Equal(got, payload) { + t.Errorf("received %d bytes, not the bytes sent", len(got)) + } + if !coalesced { + t.Errorf("no read returned more than one datagram: UDP_GRO did not coalesce") + } +} + +// TestGROSegmentSize checks the control-message parse the receive path splits by. +func TestGROSegmentSize(t *testing.T) { + for _, test := range []struct { + name string + oob []byte + want int + }{ + {name: "none"}, + {name: "segment", oob: groSegmentMessage(1200), want: 1200}, + {name: "other cmsg then segment", oob: append(udpSegmentMessage(7), groSegmentMessage(1452)...), want: 1452}, + } { + t.Run(test.name, func(t *testing.T) { + if got := groSegmentSize(test.oob); got != test.want { + t.Fatalf("groSegmentSize = %d, want %d", got, test.want) + } + }) + } +} + +// groSegmentMessage builds the UDP_GRO control message the kernel returns. +func groSegmentMessage(size int32) []byte { + const dataLen = 4 + b := make([]byte, unix.CmsgSpace(dataLen)) + header := (*unix.Cmsghdr)(unsafe.Pointer(&b[0])) + header.Level = syscall.IPPROTO_UDP + header.Type = unix.UDP_GRO + header.SetLen(unix.CmsgLen(dataLen)) + *(*int32)(unsafe.Pointer(&b[unix.CmsgSpace(0)])) = size + return b +} + +// TestIpTransportGRORecordsArrivalAddress pins the intersection of the two +// receive-side features: a wildcard-bound socket must still record the local +// address a run arrived at, even when the kernel hands the whole run to one +// read. The GRO loop is a second receive loop, so the arrival-address bookkeeping +// the ordinary loop does is not inherited -- without it a reply to this peer +// would leave from whatever address the route picks. +func TestIpTransportGRORecordsArrivalAddress(t *testing.T) { + udp, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4zero}) + if err != nil { + t.Fatal(err) + } + defer udp.Close() + recvCh := make(chan recvBatch, 16) + tr := NewIpTransport(udp, recvCh) + if !tr.gro { + t.Skip("UDP_GRO not available on this kernel") + } + if !tr.pktinfo { + t.Skip("no arrival address on this kernel") + } + + sender, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer sender.Close() + + const ( + segSize = 1200 + segs = 8 + ) + payload := make([]byte, segSize*segs) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go tr.Serve(ctx) + + dst := &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: udp.LocalAddr().(*net.UDPAddr).Port} + if _, _, err := sender.WriteMsgUDP(payload, udpSegmentMessage(segSize), dst); err != nil { + t.Fatalf("segmented write: %v", err) + } + + coalesced := false + got := 0 + deadline := time.After(10 * time.Second) + for got < len(payload) { + select { + case b := <-recvCh: + if len(batchSegments(b)) > 1 { + coalesced = true + } + got += len(b.data) + b.release() + case <-deadline: + t.Fatalf("got %d of %d bytes", got, len(payload)) + } + } + if !coalesced { + t.Skip("UDP_GRO did not coalesce: the single-datagram path is covered elsewhere") + } + + from := canonicalAddrPort(sender.LocalAddr().(*net.UDPAddr).AddrPort()) + tr.localMu.Lock() + e, ok := tr.local[from] + tr.localMu.Unlock() + if !ok { + t.Fatalf("no arrival address recorded for %s after a coalesced run", from) + } + if want := netip.MustParseAddr("127.0.0.1"); e.addr != want { + t.Fatalf("arrival address for %s = %s, want %s", from, e.addr, want) + } + if e.cmsg == nil { + t.Fatalf("no packet-info message built for %s", from) + } +} diff --git a/internal/socket/ip_gro_other.go b/internal/socket/ip_gro_other.go new file mode 100644 index 00000000..db4941ed --- /dev/null +++ b/internal/socket/ip_gro_other.go @@ -0,0 +1,11 @@ +//go:build !linux + +package socket + +import "net" + +func enableGRO(*net.UDPConn) bool { return false } + +const groOOBSize = 0 + +func groSegmentSize([]byte) int { return 0 } diff --git a/internal/socket/recvbatch.go b/internal/socket/recvbatch.go index 581abdb9..ae9ccd94 100644 --- a/internal/socket/recvbatch.go +++ b/internal/socket/recvbatch.go @@ -101,16 +101,28 @@ type RecvInfo struct { // is copied into the caller's buffer. type recvBatch struct { data []byte + stride int // segment size when data holds several datagrams; 0 = one info RecvInfo ip netip.AddrPort releaseIP bool + groBuf *[]byte releaseFn func() } +func (b recvBatch) count() uint64 { + if b.stride <= 0 || len(b.data) == 0 { + return 1 + } + return uint64((len(b.data) + b.stride - 1) / b.stride) +} + func (b recvBatch) release() { if b.releaseIP { putIPRecvBuffer(b.data) } + if b.groBuf != nil { + groRecvPool.Put(b.groBuf) + } if b.releaseFn != nil { b.releaseFn() } diff --git a/internal/socket/relay.go b/internal/socket/relay.go index 1062cccd..498e44f2 100644 --- a/internal/socket/relay.go +++ b/internal/socket/relay.go @@ -42,10 +42,7 @@ func (t *RelayTransport) Serve(ctx context.Context) { // forwardRecv drains datagrams from the actor and forwards each as a recvBatch // tagged with the relay [Addr], so [MagicConn.recvAddr] rewrites it to the -// relay mapped IPv6 ULA quic-go addresses the path by. A relay transmit may -// carry a GRO batch (segment size set); each segment is delivered as its own -// recvBatch, matching the re-batching in the Rust poll_recv -// (iroh/src/socket/transports/relay.rs:115). +// relay mapped IPv6 ULA quic-go addresses the path by. func (t *RelayTransport) forwardRecv(ctx context.Context) { for { select { @@ -60,9 +57,10 @@ func (t *RelayTransport) forwardRecv(ctx context.Context) { } } -// deliver splits dm into single datagrams (by its segment size) and forwards -// each to the recv channel. dm.Datagrams.Contents is owned by dm (the relay -// client copies on receive) and ReadFrom copies out, so segments alias it. +// deliver forwards dm as one recvBatch whose stride is the segment size, so +// ReadFrom hands quic-go one datagram at a time (the Rust poll_recv re-batches +// the same way, iroh/src/socket/transports/relay.rs:115). The pooled Contents +// buffer is released once ReadFrom has copied out the last segment. func (t *RelayTransport) deliver(ctx context.Context, dm RelayRecvDatagram) { remote := RelayAddr(dm.URL, dm.Src) b := dm.Datagrams.Contents @@ -70,21 +68,11 @@ func (t *RelayTransport) deliver(ctx context.Context, dm RelayRecvDatagram) { if dm.Datagrams.SegmentSize != 0 { stride = int(dm.Datagrams.SegmentSize) } - for { - n := min(len(b), stride) - rb := recvBatch{data: b[:n], info: RecvInfo{Remote: remote}} - if n == len(b) { - rb.releaseFn = dm.Datagrams.Release - } - select { - case t.recvCh <- rb: - case <-ctx.Done(): - return - } - b = b[n:] - if len(b) == 0 { - return - } + rb := recvBatch{data: b, stride: stride, info: RecvInfo{Remote: remote}, releaseFn: dm.Datagrams.Release} + select { + case t.recvCh <- rb: + case <-ctx.Done(): + dm.Datagrams.Release() } } diff --git a/internal/socket/transport.go b/internal/socket/transport.go index 3bae395e..5ed2db1b 100644 --- a/internal/socket/transport.go +++ b/internal/socket/transport.go @@ -47,6 +47,11 @@ type MagicConn struct { localAddr net.Addr recvCh chan recvBatch + // cur is the batch ReadFrom is draining. quic-go reads from a single + // goroutine per Transport. + cur recvBatch + curOff int + curAddr net.Addr readDeadline deadline writeDeadline deadline @@ -163,8 +168,27 @@ func (m *MagicConn) Serve(ctx context.Context) { // length and the net.Addr quic-go should associate with the path it arrived on. // For IP paths that addr is the real remote IP; for relay and custom paths it is // the synthetic mapped IPv6 ULA (port 12345). It implements net.PacketConn. +// +// ReadFrom is not safe for concurrent use: a GRO or relay batch holds several +// datagrams, and successive calls hand them out one at a time from a cursor. +// Only the quic-go Transport reads from a MagicConn, from its single listen +// goroutine. func (m *MagicConn) ReadFrom(p []byte) (int, net.Addr, error) { for { + if m.curOff < len(m.cur.data) { + seg := m.cur.data[m.curOff:] + if m.cur.stride > 0 && len(seg) > m.cur.stride { + seg = seg[:m.cur.stride] + } + m.curOff += len(seg) + n := copy(p, seg) + addr := m.curAddr + if m.curOff >= len(m.cur.data) { + m.cur.release() + m.cur, m.curAddr = recvBatch{}, nil + } + return n, addr, nil + } select { case b := <-m.recvCh: addr, ok := m.recvBatchAddr(b) @@ -174,10 +198,14 @@ func (m *MagicConn) ReadFrom(p []byte) (int, net.Addr, error) { // quic-go. Drop and keep reading. continue } - m.recordRecv(b.recvAddr()) - n := copy(p, b.data) - b.release() - return n, addr, nil + m.recordRecv(b.recvAddr(), b.count()) + if len(b.data) == 0 { + // Nothing to hand out, and an empty cursor would be + // overwritten unreleased by the next batch. + b.release() + continue + } + m.cur, m.curOff, m.curAddr = b, 0, addr case <-m.readDeadline.wait(): return 0, nil, timeoutError{} } @@ -446,20 +474,20 @@ func (m *MagicConn) SendAddr(addr Addr, p []byte) bool { return m.sendAddr(addr, p) } -func (m *MagicConn) recordRecv(addr Addr) { - m.metrics.recvDatagrams.Add(1) +func (m *MagicConn) recordRecv(addr Addr, n uint64) { + m.metrics.recvDatagrams.Add(n) switch addr.Kind() { case AddrIP: ap, _ := addr.IP() if ap.Addr().Is4() { - m.metrics.ipv4Recv.Add(1) + m.metrics.ipv4Recv.Add(n) } else { - m.metrics.ipv6Recv.Add(1) + m.metrics.ipv6Recv.Add(n) } case AddrRelay: - m.metrics.relayRecv.Add(1) + m.metrics.relayRecv.Add(n) case AddrCustom: - m.metrics.customRecv.Add(1) + m.metrics.customRecv.Add(n) } } diff --git a/iroh/defaults.go b/iroh/defaults.go index 8eb666d0..6bd48bfd 100644 --- a/iroh/defaults.go +++ b/iroh/defaults.go @@ -36,5 +36,9 @@ const ( const ConnectTimeout = 10 * time.Second // dialAttemptTimeout bounds how long a dial waits for an unproven target's -// handshake before trying the next one. See Endpoint.connectEarly. +// handshake before giving up on it. See Endpoint.connectEarly. const dialAttemptTimeout = 3 * time.Second + +// dialAttemptDelay staggers the handshakes Endpoint.dialAny starts across a +// peer's dial targets. +const dialAttemptDelay = 250 * time.Millisecond diff --git a/iroh/endpoint.go b/iroh/endpoint.go index 42f8dfd3..a4f44fe2 100644 --- a/iroh/endpoint.go +++ b/iroh/endpoint.go @@ -1286,9 +1286,11 @@ var ErrHandshakeRejected = errors.New("iroh: handshake rejected by hook") var ErrConnClosedDuringHandshake = errors.New("iroh: connection closed during handshake") // Connect dials the endpoint identified by addr and negotiates alpn, returning -// an established [Conn]. It tries the direct IP addresses in addr in order, then -// (if relays are enabled) the relay URLs in addr. A relay path carries the QUIC -// handshake over a relay mapped address that routes through the relay transport. +// an established [Conn]. It starts a handshake on the direct IP addresses in +// addr and then (if relays are enabled) on the relay URLs, in that order but +// staggered rather than one after the other, and takes the first that completes. +// A relay path carries the QUIC handshake over a relay mapped address that +// routes through the relay transport. // // When addr carries no usable address at all - a dial by endpoint ID alone - // Connect resolves one through the services given to [WithAddressLookup] and @@ -1394,36 +1396,103 @@ func (e *Endpoint) connectEarly(ctx context.Context, addr netaddr.EndpointAddr, // With a resumable ticket DialEarly returns at the 0-RTT window, before any // packet from the peer, so success says nothing about whether the target - // answers. Handshake completion is the first real evidence, needed for - // every target except one already proven or with nothing left to try. - sel, haveSel := e.goodTarget(addr.ID) + // answers. Handshake completion is the first real evidence, needed for every + // target except one the path selector has already proven. var firstErr error - for i, target := range dials { + if target, ok := e.provenTarget(addr.ID, dials); ok { qc, err := e.transport.DialEarly(ctx, target, clientTLS, e.quicConf) - if err != nil { - if firstErr == nil { - firstErr = err - } - continue + if err == nil { + return &Connecting{ep: e, qc: qc, remoteID: addr.ID, addr: addr, alpn: alpn}, nil + } + firstErr = err + } + qc, err := e.dialAny(ctx, dials, clientTLS) + if err != nil { + if firstErr == nil { + firstErr = err } - proven := haveSel && e.sock.PathAddr(addr.ID, target).String() == sel - if !proven && i < len(dials)-1 { - select { - case <-qc.HandshakeComplete(): - case <-ctx.Done(): - qc.CloseWithError(0, "") - return nil, fmt.Errorf("iroh: connect to %s: %w", addr.ID, ctx.Err()) - case <-time.After(dialAttemptTimeout): - qc.CloseWithError(0, "") - if firstErr == nil { - firstErr = fmt.Errorf("dial %s: no handshake within %v", target, dialAttemptTimeout) + return nil, fmt.Errorf("iroh: connect to %s: %w", addr.ID, tlsHandshakeFailure(firstErr)) + } + return &Connecting{ep: e, qc: qc, remoteID: addr.ID, addr: addr, alpn: alpn}, nil +} + +// provenTarget returns the one of targets that the path selector currently +// prefers for id, if any. A dial to it needs no further evidence that id +// answers there. +func (e *Endpoint) provenTarget(id key.EndpointID, targets []net.Addr) (net.Addr, bool) { + sel, ok := e.goodTarget(id) + if !ok { + return nil, false + } + for _, target := range targets { + if e.sock.PathAddr(id, target).String() == sel { + return target, true + } + } + return nil, false +} + +// dialAny starts a handshake with each target dialAttemptDelay apart and returns +// the first whose handshake completes, so unreachable direct addresses delay the +// relay path by that step rather than by dialAttemptTimeout each. Losing +// handshakes are cancelled or closed; the peer may briefly accept one. +func (e *Endpoint) dialAny(ctx context.Context, targets []net.Addr, clientTLS *itls.Config) (*quic.Conn, error) { + if len(targets) == 1 { + // Nothing left to race against, so there is no winner to pick and no + // reason to wait past the 0-RTT window. + return e.transport.DialEarly(ctx, targets[0], clientTLS, e.quicConf) + } + type result struct { + qc *quic.Conn + err error + } + ctx, cancel := context.WithCancel(ctx) + defer cancel() + results := make(chan result, len(targets)) + for i, target := range targets { + go func() { + if i > 0 { + select { + case <-time.After(time.Duration(i) * dialAttemptDelay): + case <-ctx.Done(): + results <- result{nil, ctx.Err()} + return + } + } + qc, err := e.transport.DialEarly(ctx, target, clientTLS.Clone(), e.quicConf) + if err == nil { + select { + case <-qc.HandshakeComplete(): + case <-ctx.Done(): + qc.CloseWithError(0, "") + qc, err = nil, ctx.Err() + case <-time.After(dialAttemptTimeout): + qc.CloseWithError(0, "") + qc, err = nil, fmt.Errorf("dial %s: no handshake within %v", target, dialAttemptTimeout) } - continue } + results <- result{qc, err} + }() + } + var firstErr error + for got := 1; got <= len(targets); got++ { + r := <-results + if r.err == nil { + cancel() + go func() { + for range len(targets) - got { + if late := <-results; late.qc != nil { + late.qc.CloseWithError(0, "") + } + } + }() + return r.qc, nil + } + if firstErr == nil { + firstErr = r.err } - return &Connecting{ep: e, qc: qc, remoteID: addr.ID, addr: addr, alpn: alpn}, nil } - return nil, fmt.Errorf("iroh: connect to %s: %w", addr.ID, tlsHandshakeFailure(firstErr)) + return nil, firstErr } // Dial dials addr, negotiates alpn, opens a bidirectional stream, and returns it diff --git a/iroh/relay_echo_test.go b/iroh/relay_echo_test.go index 417605ca..606d828b 100644 --- a/iroh/relay_echo_test.go +++ b/iroh/relay_echo_test.go @@ -4,6 +4,7 @@ import ( "context" "io" "net/http/httptest" + "net/netip" "testing" "time" @@ -147,3 +148,53 @@ func TestRelayOnlyEcho(t *testing.T) { t.Errorf("server saw client id %s, want %s", res.peer, client.ID()) } } + +// TestConnectRacesDialTargets checks that unreachable direct addresses do not +// delay the relay path by a handshake timeout each. +func TestConnectRacesDialTargets(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + srv := newEchoRelayServer(t) + relayURL := srv.url(t) + mode := relay.ModeCustom(relay.MapFromURLs(relayURL)) + const alpn = "iroh-race/0" + + server, err := Bind(ctx, WithALPNs(alpn), WithRelayMode(mode), WithBindAddr(netip.MustParseAddrPort("127.0.0.1:0"))) + if err != nil { + t.Fatal(err) + } + defer server.Shutdown(ctx) + client, err := Bind(ctx, WithRelayMode(mode), WithBindAddr(netip.MustParseAddrPort("127.0.0.1:0"))) + if err != nil { + t.Fatal(err) + } + defer client.Shutdown(ctx) + if err := server.Online(ctx); err != nil { + t.Fatal(err) + } + if err := client.Online(ctx); err != nil { + t.Fatal(err) + } + go func() { + if conn, err := server.Accept(ctx); err == nil { + <-ctx.Done() + conn.CloseWithError(0, "") + } + }() + + // Three blackholed direct addresses (TEST-NET-1), then the relay. + addr := netaddr.NewEndpointAddr(server.ID()). + WithIP(netip.MustParseAddrPort("192.0.2.1:7")). + WithIP(netip.MustParseAddrPort("192.0.2.2:7")). + WithIP(netip.MustParseAddrPort("192.0.2.3:7")). + WithRelayURL(relayURL) + start := time.Now() + conn, err := client.Connect(ctx, addr, alpn) + if err != nil { + t.Fatalf("connect: %v", err) + } + defer conn.CloseWithError(0, "") + if d := time.Since(start); d > 3*time.Second { + t.Fatalf("connect took %v, direct targets were tried sequentially", d) + } +} diff --git a/iroh/zerortt_interception_test.go b/iroh/zerortt_interception_test.go index 720341eb..65c9e638 100644 --- a/iroh/zerortt_interception_test.go +++ b/iroh/zerortt_interception_test.go @@ -436,8 +436,9 @@ func TestConnectEarlyUnreachableFirstTarget(t *testing.T) { // TestConnectEarlyFallsThroughUnprovenTarget covers the dial path when no // proven target is available: the hint recorded by an earlier connection is -// dropped when the remote is evicted, so a resumed dial must wait for evidence -// on the unreachable first target and fall through to the next one. +// dropped when the remote is evicted, so a resumed dial has to find evidence +// that a target answers before it returns one. With no proven target the dial +// races them, and the blackholed address must lose to the reachable one. func TestConnectEarlyFallsThroughUnprovenTarget(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() @@ -487,13 +488,20 @@ func TestConnectEarlyFallsThroughUnprovenTarget(t *testing.T) { conn2, _ := c2.Into0RTT() defer conn2.CloseWithError(0, "") + // A resumed handshake is no evidence that a target answers, so the dial + // must have landed on the reachable address, not on the blackholed one. + got, err := netip.ParseAddrPort(conn2.RemoteAddr().String()) + if err != nil { + t.Fatalf("remote addr %q: %v", conn2.RemoteAddr(), err) + } + want := server.LocalAddr() + if got.Addr().Unmap() != want.Addr().Unmap() || got.Port() != want.Port() { + t.Fatalf("dial returned %v after %v, want %v: it committed to an unproven address", got, time.Since(start), want) + } select { case <-conn2.HandshakeComplete(): case <-ctx.Done(): t.Fatalf("warm handshake did not complete: %v", ctx.Err()) } - if waited := time.Since(start); waited < dialAttemptTimeout { - t.Errorf("dial returned after %v, want at least dialAttemptTimeout (%v): the unproven first target was not waited on", waited, dialAttemptTimeout) - } echo(t, ctx, conn2, "warm") }