diff --git a/internal/socket/export_test.go b/internal/socket/export_test.go new file mode 100644 index 00000000..4d12cf23 --- /dev/null +++ b/internal/socket/export_test.go @@ -0,0 +1,9 @@ +package socket + +// IPv4ArrivalSupported reports whether this platform reports the IPv4 +// destination of a received datagram in a form parsePacketInfo understands. +// Tests that expect an IPv4 reply to leave from the arrival address skip +// where it does not: the BSDs report the destination through IP_RECVDSTADDR, +// which cannot be paired with an IPv4 source on send, so a reply there leaves +// from the address the kernel picks, as it did before the table. +const IPv4ArrivalSupported = sysIPPktinfo >= 0 diff --git a/internal/socket/ip.go b/internal/socket/ip.go index 21286c11..1f8e5acd 100644 --- a/internal/socket/ip.go +++ b/internal/socket/ip.go @@ -5,6 +5,10 @@ import ( "errors" "net" "net/netip" + "sync" + + "golang.org/x/net/ipv4" + "golang.org/x/net/ipv6" ) // maxDatagramSize bounds a single read from the UDP socket. QUIC packets never @@ -22,15 +26,162 @@ var ipRecvPool = make(chan []byte, 1024) type IpTransport struct { conn *net.UDPConn recvCh chan<- recvBatch + + // 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 + // reply by route, which on a multi-homed host, or one whose peer is + // reached over a bridge or a container network, need not be the address + // the peer sent to. A QUIC peer then receives packets from an address it + // never dialed and fails path validation. The transport records the local + // address each remote last reached and answers from it, as quinn does + // for the Rust iroh through quinn-udp's RecvMeta::dst_ip and + // Transmit::src_ip. + pktinfo bool + localMu sync.Mutex + local map[netip.AddrPort]*localEntry + // heard orders entries by the last datagram heard from each remote, + // circularly through this sentinel: heard.next is newest, heard.prev + // oldest and first to be evicted. + heard localEntry +} + +// localAddr is the local address a remote's datagrams arrive at. +type localAddr struct { + addr netip.Addr + ifIndex int +} + +// localEntry is a remote's entry in the arrival-address table. cmsg is the +// packet-info message that makes a reply leave from the address; it is +// shared with senders and never written after it is built. +type localEntry struct { + remote netip.AddrPort + localAddr + cmsg []byte + prev, next *localEntry } +// maxLocalAddrs bounds the arrival-address table. A full table drops the +// remote heard from longest ago; a reply to a peer not in the table leaves +// from the address the kernel picks, as before the table. +const maxLocalAddrs = 4096 + // 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 { - return &IpTransport{conn: conn, recvCh: recvCh} + t := &IpTransport{conn: conn, recvCh: recvCh} + if la, ok := conn.LocalAddr().(*net.UDPAddr); ok && la.IP.IsUnspecified() { + t.pktinfo = enablePacketInfo(conn) + if t.pktinfo { + t.local = make(map[netip.AddrPort]*localEntry) + t.heard.prev, t.heard.next = &t.heard, &t.heard + } + } + return t } +// enablePacketInfo asks the kernel for the destination address of received +// datagrams. Either family may fail: an IPv4 socket has no IPv6 options, and +// some platforms have neither. It reports whether at least one family took. +func enablePacketInfo(conn *net.UDPConn) bool { + ok4 := ipv4.NewPacketConn(conn).SetControlMessage(ipv4.FlagDst|ipv4.FlagInterface, true) == nil + ok6 := ipv6.NewPacketConn(conn).SetControlMessage(ipv6.FlagDst|ipv6.FlagInterface, true) == nil + return ok4 || ok6 +} + +// recordLocal notes the local address remote's datagram arrived at. +func (t *IpTransport) recordLocal(remote netip.AddrPort, oob []byte) { + if la, ok := parsePacketInfo(oob); ok { + t.record(remote, la) + } +} + +// record notes la as remote's arrival address and remote as the most +// recently heard. The control message is rebuilt only when la changes. +func (t *IpTransport) record(remote netip.AddrPort, la localAddr) { + t.localMu.Lock() + defer t.localMu.Unlock() + if e, ok := t.local[remote]; ok { + if e.localAddr != la { + e.localAddr = la + e.cmsg = packetInfoMessage(remote, la) + } + t.unlink(e) + t.linkFront(e) + return + } + if len(t.local) >= maxLocalAddrs { + oldest := t.heard.prev + t.unlink(oldest) + delete(t.local, oldest.remote) + } + e := &localEntry{remote: remote, localAddr: la, cmsg: packetInfoMessage(remote, la)} + t.local[remote] = e + t.linkFront(e) +} + +// linkFront puts e at the head of the recency list. Caller holds localMu. +func (t *IpTransport) linkFront(e *localEntry) { + e.prev, e.next = &t.heard, t.heard.next + e.prev.next, e.next.prev = e, e +} + +// unlink removes e from the recency list. Caller holds localMu. +func (t *IpTransport) unlink(e *localEntry) { + e.prev.next, e.next.prev = e.next, e.prev + e.prev, e.next = nil, nil +} + +// forgetLocal drops remote's arrival address after the kernel refused it. +func (t *IpTransport) forgetLocal(remote netip.AddrPort) { + t.localMu.Lock() + if e, ok := t.local[remote]; ok { + t.unlink(e) + delete(t.local, remote) + } + t.localMu.Unlock() +} + +// packetInfoFor returns the control message that makes a datagram to remote +// leave from the address remote last reached, or nil when none is known. The +// message is shared and must not be modified. +func (t *IpTransport) packetInfoFor(remote netip.AddrPort) []byte { + if !t.pktinfo { + return nil + } + var cmsg []byte + t.localMu.Lock() + if e, ok := t.local[remote]; ok { + cmsg = e.cmsg + } + t.localMu.Unlock() + return cmsg +} + +// packetInfoMessage builds the packet-info control message naming la as the +// source of datagrams to remote, or nil when it cannot be. The source must +// be of the destination's family: IPv4 traffic on a dual-stack socket +// carries IP-level control messages in both directions, so an IPv4 remote +// gets an IPv4 message whatever the socket. +func packetInfoMessage(remote netip.AddrPort, la localAddr) []byte { + if remote.Addr().Is4() { + if !la.addr.Is4() { + return nil + } + return (&ipv4.ControlMessage{Src: la.addr.AsSlice()}).Marshal() + } + if !la.addr.Is6() || la.addr.Is4In6() { + return nil + } + return (&ipv6.ControlMessage{Src: la.addr.AsSlice(), IfIndex: la.ifIndex}).Marshal() +} + +// maxControlSize bounds the control messages read with a datagram: a +// packet-info message of either family plus headroom. +const maxControlSize = 128 + // LocalAddr returns the bound local address of the underlying socket. func (t *IpTransport) LocalAddr() net.Addr { return t.conn.LocalAddr() } @@ -40,12 +191,25 @@ 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) { + var oob []byte + if t.pktinfo { + oob = make([]byte, maxControlSize) + } for { if ctx.Err() != nil { return } buf := getIPRecvBuffer() - n, ap, err := t.conn.ReadFromUDPAddrPort(buf) + var ( + n, oobn int + ap netip.AddrPort + err error + ) + if t.pktinfo { + n, oobn, _, ap, err = t.conn.ReadMsgUDPAddrPort(buf, oob) + } else { + n, ap, err = t.conn.ReadFromUDPAddrPort(buf) + } if err != nil { putIPRecvBuffer(buf) if errors.Is(err, net.ErrClosed) || ctx.Err() != nil { @@ -64,6 +228,9 @@ func (t *IpTransport) Serve(ctx context.Context) { // The transport address is internal to iroh and is always the canonical // (unmapped) form. iroh/src/socket/transports/ip.rs:219. cap := canonicalAddrPort(ap) + if t.pktinfo { + t.recordLocal(cap, oob[:oobn]) + } b := recvBatch{data: buf[:n], ip: cap, releaseIP: true} if !t.enqueue(ctx, b) { return @@ -107,14 +274,37 @@ func putIPRecvBuffer(buf []byte) { // send writes p to the IP destination dst. The destination is canonicalized so // an IPv4-mapped IPv6 address is sent as plain IPv4, matching -// iroh/src/socket/transports/ip.rs:310 canonical_addr. It reports the number of -// bytes written. +// iroh/src/socket/transports/ip.rs:310 canonical_addr. When dst has reached +// this socket before, the datagram leaves from the address it arrived at. It +// reports the number of bytes written. func (t *IpTransport) send(p []byte, dst netip.AddrPort) (int, error) { dst = canonicalAddrPort(dst) + if oob := t.packetInfoFor(dst); oob != nil { + n, _, err := t.conn.WriteMsgUDPAddrPort(p, oob, dst) + if err == nil { + return n, nil + } + // The kernel refused the source (an address that has since gone, + // or a platform quirk): fall back to its own choice, and stop asking + // until the peer reaches us again. + t.forgetLocal(dst) + } n, err := t.conn.WriteToUDPAddrPort(p, dst) return n, err } +// withPacketInfo prepends the packet-info control message for dst, when one +// is known, to oob. The result is built in buf, grown if too small; oob and +// the shared message are not written to. +func (t *IpTransport) withPacketInfo(dst netip.AddrPort, oob, buf []byte) []byte { + pi := t.packetInfoFor(dst) + if pi == nil { + return oob + } + buf = append(buf[:0], pi...) + return append(buf, oob...) +} + func canonicalAddrPort(ap netip.AddrPort) netip.AddrPort { addr := ap.Addr() if !addr.Is4In6() { diff --git a/internal/socket/ip_pktinfo_bench_test.go b/internal/socket/ip_pktinfo_bench_test.go new file mode 100644 index 00000000..3579088c --- /dev/null +++ b/internal/socket/ip_pktinfo_bench_test.go @@ -0,0 +1,132 @@ +package socket + +import ( + "context" + "net" + "net/netip" + "testing" + "time" +) + +// pktinfoBenchConn returns a wildcard-bound MagicConn that has received a +// datagram from a loopback client, so the client's arrival address is +// recorded, with the client and its address. +func pktinfoBenchConn(b *testing.B) (*MagicConn, *net.UDPConn, netip.AddrPort) { + b.Helper() + udp, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4zero}) + if err != nil { + b.Fatal(err) + } + b.Cleanup(func() { udp.Close() }) + m := NewMagicConn(NewSocket(), udp) + ctx, cancel := context.WithCancel(context.Background()) + b.Cleanup(cancel) + go m.Serve(ctx) + b.Cleanup(func() { m.Close() }) + if !m.transports.ip.pktinfo { + b.Skip("packet info not available on a wildcard socket") + } + client, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + b.Fatal(err) + } + b.Cleanup(func() { client.Close() }) + server := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), uint16(udp.LocalAddr().(*net.UDPAddr).Port)) + if _, err := client.WriteToUDPAddrPort([]byte("ping"), server); err != nil { + b.Fatal(err) + } + m.SetReadDeadline(time.Now().Add(2 * time.Second)) + buf := make([]byte, 64) + if _, _, err := m.ReadFrom(buf); err != nil { + b.Fatal(err) + } + m.SetReadDeadline(time.Time{}) + return m, client, client.LocalAddr().(*net.UDPAddr).AddrPort() +} + +func BenchmarkIpTransportSend(b *testing.B) { + m, _, clientAddr := pktinfoBenchConn(b) + payload := make([]byte, 1200) + + b.Run("recorded", func(b *testing.B) { + dst := net.UDPAddrFromAddrPort(clientAddr) + b.ReportAllocs() + for b.Loop() { + m.WriteTo(payload, dst) + } + b.StopTimer() + if n := m.metrics.blackholed.Load(); n != 0 { + b.Fatalf("%d sends blackholed", n) + } + if m.transports.ip.packetInfoFor(clientAddr) == nil { + b.Fatal("arrival address was dropped during the benchmark") + } + }) + + b.Run("unrecorded", func(b *testing.B) { + other, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + b.Fatal(err) + } + defer other.Close() + dst := other.LocalAddr().(*net.UDPAddr) + b.ReportAllocs() + for b.Loop() { + m.WriteTo(payload, dst) + } + b.StopTimer() + if n := m.metrics.blackholed.Load(); n != 0 { + b.Fatalf("%d sends blackholed", n) + } + }) +} + +// BenchmarkIpTransportRecordLocal measures recording an arrival address from +// control messages the kernel delivered. +func BenchmarkIpTransportRecordLocal(b *testing.B) { + udp, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4zero}) + if err != nil { + b.Fatal(err) + } + defer udp.Close() + t := NewIpTransport(udp, make(chan recvBatch, 1)) + if !t.pktinfo { + b.Skip("packet info not available on a wildcard socket") + } + client, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + b.Fatal(err) + } + defer client.Close() + server := netip.AddrPortFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), uint16(udp.LocalAddr().(*net.UDPAddr).Port)) + if _, err := client.WriteToUDPAddrPort([]byte("ping"), server); err != nil { + b.Fatal(err) + } + buf, oob := make([]byte, 64), make([]byte, maxControlSize) + udp.SetReadDeadline(time.Now().Add(2 * time.Second)) + _, oobn, _, from, err := udp.ReadMsgUDPAddrPort(buf, oob) + if err != nil { + b.Fatal(err) + } + oob = oob[:oobn] + la, ok := parsePacketInfo(oob) + if !ok || la.addr != server.Addr() { + b.Fatalf("parsePacketInfo(%d bytes) = %+v, %v; want %s", oobn, la, ok, server.Addr()) + } + b.ReportAllocs() + for b.Loop() { + t.recordLocal(from, oob) + } +} + +// BenchmarkIpTransportWithPacketInfo measures the control-message merge on +// the segmented (GSO) send path. +func BenchmarkIpTransportWithPacketInfo(b *testing.B) { + m, _, clientAddr := pktinfoBenchConn(b) + oob := make([]byte, 24) // stands in for a UDP_SEGMENT control message + b.ReportAllocs() + for b.Loop() { + var cbuf [maxControlSize]byte + _ = m.transports.ip.withPacketInfo(clientAddr, oob, cbuf[:0]) + } +} diff --git a/internal/socket/ip_pktinfo_linux_test.go b/internal/socket/ip_pktinfo_linux_test.go new file mode 100644 index 00000000..3b347479 --- /dev/null +++ b/internal/socket/ip_pktinfo_linux_test.go @@ -0,0 +1,141 @@ +//go:build linux + +package socket + +import ( + "context" + "net" + "net/netip" + "testing" + "time" +) + +func testLANIPv4(t *testing.T) netip.Addr { + t.Helper() + ifaces, err := net.Interfaces() + if err != nil { + t.Skipf("interfaces: %v", err) + } + for _, ifi := range ifaces { + if ifi.Flags&net.FlagUp == 0 || ifi.Flags&net.FlagLoopback != 0 { + continue + } + addrs, _ := ifi.Addrs() + for _, a := range addrs { + if pfx, err := netip.ParsePrefix(a.String()); err == nil && pfx.Addr().Unmap().Is4() { + return pfx.Addr().Unmap() + } + } + } + t.Skip("no non-loopback IPv4 interface address") + return netip.Addr{} +} + +// TestMagicConnWriteMsgUDPKeepsArrivalAddress checks that a segmented (GSO) +// send to a peer that reached a wildcard socket carries the packet-info +// control message next to UDP_SEGMENT, so every segment leaves from the +// address the peer sent to. +func TestMagicConnWriteMsgUDPKeepsArrivalAddress(t *testing.T) { + lan := testLANIPv4(t) + udp, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4zero}) + if err != nil { + t.Fatal(err) + } + defer udp.Close() + sock := NewSocket() + m := NewMagicConn(sock, udp) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go m.Serve(ctx) + defer m.Close() + if !m.transports.ip.pktinfo { + t.Fatal("packet info not enabled on a wildcard socket") + } + + client, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + server := netip.AddrPortFrom(lan, uint16(udp.LocalAddr().(*net.UDPAddr).Port)) + if _, err := client.WriteToUDPAddrPort([]byte("ping"), server); err != nil { + t.Skipf("cannot send from loopback to %s: %v", server, err) + } + m.SetReadDeadline(time.Now().Add(2 * time.Second)) + buf := make([]byte, 64) + if _, _, err := m.ReadFrom(buf); err != nil { + t.Skipf("no datagram from loopback to %s arrived: %v", server, err) + } + + payload := []byte("abcdefg") + if _, _, err := m.WriteMsgUDP(payload, udpSegmentMessage(3), client.LocalAddr().(*net.UDPAddr)); err != nil { + t.Fatal(err) + } + var got []string + for range 3 { + client.SetReadDeadline(time.Now().Add(2 * time.Second)) + n, from, err := client.ReadFromUDPAddrPort(buf) + if err != nil { + t.Fatalf("segment %d: %v", len(got), err) + } + if from != server { + t.Fatalf("segment %d came from %s, want %s", len(got), from, server) + } + got = append(got, string(buf[:n])) + } + if want := []string{"abc", "def", "g"}; len(got) != 3 || got[0] != want[0] || got[1] != want[1] || got[2] != want[2] { + t.Fatalf("segments %q, want %q", got, want) + } +} + +// TestMagicConnWriteMsgUDPFallsBackWhenSourceRefused checks that a send whose +// recorded arrival address the kernel refuses is retried without it and the +// address forgotten. IpTransport.send has done this since the table existed; +// WriteMsgUDP needs it too, and more: every ECN-marked packet reaches the +// wire through it, so a stale entry there blackholes a peer until it sends +// to us again. +func TestMagicConnWriteMsgUDPFallsBackWhenSourceRefused(t *testing.T) { + udp, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4zero}) + if err != nil { + t.Fatal(err) + } + defer udp.Close() + sock := NewSocket() + m := NewMagicConn(sock, udp) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go m.Serve(ctx) + defer m.Close() + if !m.transports.ip.pktinfo { + t.Fatal("packet info not enabled on a wildcard socket") + } + + client, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + peer := canonicalAddrPort(client.LocalAddr().(*net.UDPAddr).AddrPort()) + + // An address this host does not hold, standing in for one that has gone. + m.transports.ip.record(peer, localAddr{addr: netip.MustParseAddr("192.0.2.1")}) + if m.transports.ip.packetInfoFor(peer) == nil { + t.Fatal("arrival address not recorded") + } + + if _, _, err := m.WriteMsgUDP([]byte("pong"), nil, client.LocalAddr().(*net.UDPAddr)); err != nil { + t.Fatalf("WriteMsgUDP: %v", err) + } + client.SetReadDeadline(time.Now().Add(2 * time.Second)) + buf := make([]byte, 64) + n, _, err := client.ReadFromUDPAddrPort(buf) + if err != nil { + t.Fatalf("the datagram never left: %v", err) + } + if string(buf[:n]) != "pong" { + t.Fatalf("received %q", buf[:n]) + } + if m.transports.ip.packetInfoFor(peer) != nil { + t.Error("the refused arrival address is still recorded") + } +} diff --git a/internal/socket/ip_pktinfo_lru_test.go b/internal/socket/ip_pktinfo_lru_test.go new file mode 100644 index 00000000..5ec5738a --- /dev/null +++ b/internal/socket/ip_pktinfo_lru_test.go @@ -0,0 +1,99 @@ +package socket + +import ( + "bytes" + "net" + "net/netip" + "testing" +) + +// newPktinfoTransport returns an IpTransport on a wildcard socket, skipping +// when the platform delivers no packet info. +func newPktinfoTransport(t *testing.T) *IpTransport { + t.Helper() + udp, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4zero}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { udp.Close() }) + tr := NewIpTransport(udp, make(chan recvBatch, 1)) + if !tr.pktinfo { + t.Skip("packet info not available on a wildcard socket") + } + return tr +} + +// remoteN returns a distinct remote address for n. +func remoteN(n int) netip.AddrPort { + return netip.AddrPortFrom(netip.AddrFrom4([4]byte{10, byte(n >> 16), byte(n >> 8), byte(n)}), 4434) +} + +var arrival = localAddr{addr: netip.AddrFrom4([4]byte{192, 0, 2, 1})} + +// TestIpTransportEvictsLeastRecentlyHeardRemote checks that a full table +// drops only the remote heard from longest ago. +func TestIpTransportEvictsLeastRecentlyHeardRemote(t *testing.T) { + tr := newPktinfoTransport(t) + for n := 1; n <= maxLocalAddrs; n++ { + tr.record(remoteN(n), arrival) + } + tr.record(remoteN(1), arrival) // heard from again: now the most recent + tr.record(remoteN(maxLocalAddrs+1), arrival) + + if got := len(tr.local); got != maxLocalAddrs { + t.Errorf("table holds %d remotes, want %d", got, maxLocalAddrs) + } + for _, tc := range []struct { + n int + want bool + }{ + {1, true}, // refreshed after the table filled + {2, false}, // heard from longest ago + {3, true}, // one eviction made room for one newcomer + {maxLocalAddrs, true}, // the last of the first fill + {maxLocalAddrs + 1, true}, // the newcomer + } { + if got := tr.packetInfoFor(remoteN(tc.n)) != nil; got != tc.want { + t.Errorf("remote %d recorded = %v, want %v", tc.n, got, tc.want) + } + } +} + +// TestIpTransportForgottenRemoteLeavesRecencyOrder checks that a forgotten +// remote frees its place and does not absorb a later eviction. +func TestIpTransportForgottenRemoteLeavesRecencyOrder(t *testing.T) { + tr := newPktinfoTransport(t) + for n := 1; n <= maxLocalAddrs; n++ { + tr.record(remoteN(n), arrival) + } + tr.forgetLocal(remoteN(1)) + tr.record(remoteN(maxLocalAddrs+1), arrival) // takes the freed place + tr.record(remoteN(maxLocalAddrs+2), arrival) // evicts remote 2, the oldest left + + if got := len(tr.local); got != maxLocalAddrs { + t.Errorf("table holds %d remotes, want %d", got, maxLocalAddrs) + } + if tr.packetInfoFor(remoteN(2)) != nil { + t.Error("remote 2 still recorded: the forgotten remote 1 absorbed the eviction") + } + if tr.packetInfoFor(remoteN(3)) == nil { + t.Error("remote 3 was evicted") + } +} + +// TestIpTransportRecordReplacesChangedArrivalAddress checks that a remote +// reaching the socket at a new local address gets a new control message and +// keeps a single entry. +func TestIpTransportRecordReplacesChangedArrivalAddress(t *testing.T) { + tr := newPktinfoTransport(t) + r := remoteN(1) + moved := localAddr{addr: netip.AddrFrom4([4]byte{192, 0, 2, 2})} + tr.record(r, arrival) + tr.record(r, moved) + if got, want := tr.packetInfoFor(r), packetInfoMessage(r, moved); !bytes.Equal(got, want) { + t.Errorf("control message % x, want % x", got, want) + } + if got := len(tr.local); got != 1 { + t.Errorf("table holds %d remotes, want 1", got) + } +} diff --git a/internal/socket/ip_pktinfo_other.go b/internal/socket/ip_pktinfo_other.go new file mode 100644 index 00000000..1eea2aaa --- /dev/null +++ b/internal/socket/ip_pktinfo_other.go @@ -0,0 +1,10 @@ +//go:build !unix + +package socket + +// sysIPPktinfo matches nothing: there is no packet info to parse here. +const sysIPPktinfo = -1 + +// parsePacketInfo reports no arrival address: x/net delivers no packet info +// here, so the transport never asks for one. +func parsePacketInfo([]byte) (localAddr, bool) { return localAddr{}, false } diff --git a/internal/socket/ip_pktinfo_sys_other.go b/internal/socket/ip_pktinfo_sys_other.go new file mode 100644 index 00000000..8e8b392a --- /dev/null +++ b/internal/socket/ip_pktinfo_sys_other.go @@ -0,0 +1,8 @@ +//go:build unix && !linux && !darwin + +package socket + +// sysIPPktinfo matches nothing: the BSDs report an IPv4 destination through +// IP_RECVDSTADDR, which x/net cannot pair with an IPv4 source on send. IPv6 +// packet info (RFC 3542) is recognized. +const sysIPPktinfo = -1 diff --git a/internal/socket/ip_pktinfo_sys_pktinfo.go b/internal/socket/ip_pktinfo_sys_pktinfo.go new file mode 100644 index 00000000..976e2a21 --- /dev/null +++ b/internal/socket/ip_pktinfo_sys_pktinfo.go @@ -0,0 +1,8 @@ +//go:build linux || darwin + +package socket + +import "golang.org/x/sys/unix" + +// sysIPPktinfo is the control message type carrying struct in_pktinfo. +const sysIPPktinfo = unix.IP_PKTINFO diff --git a/internal/socket/ip_pktinfo_test.go b/internal/socket/ip_pktinfo_test.go new file mode 100644 index 00000000..539e9c7a --- /dev/null +++ b/internal/socket/ip_pktinfo_test.go @@ -0,0 +1,104 @@ +package socket_test + +import ( + "context" + "net" + "net/netip" + "testing" + "time" + + "github.com/tmc/go-iroh/internal/socket" +) + +// lanIPv4 returns a unicast IPv4 address of a non-loopback interface that is +// up, or skips the test. +func lanIPv4(t *testing.T) netip.Addr { + t.Helper() + ifaces, err := net.Interfaces() + if err != nil { + t.Skipf("interfaces: %v", err) + } + for _, ifi := range ifaces { + if ifi.Flags&net.FlagUp == 0 || ifi.Flags&net.FlagLoopback != 0 { + continue + } + addrs, err := ifi.Addrs() + if err != nil { + continue + } + for _, a := range addrs { + pfx, err := netip.ParsePrefix(a.String()) + if err != nil { + continue + } + ip := pfx.Addr().Unmap() + if ip.Is4() && ip.IsGlobalUnicast() || ip.Is4() && ip.IsPrivate() { + return ip + } + } + } + t.Skip("no non-loopback IPv4 interface address") + return netip.Addr{} +} + +// TestIpTransportRepliesFromArrivalAddress checks that a wildcard-bound +// socket answers from the local address a datagram arrived on, not from the +// address the kernel picks by route. The client sends from 127.0.0.1 to the +// host's LAN address; by route the reply would leave from 127.0.0.1, and a +// QUIC peer would then see a packet from an address it never dialed. +func TestIpTransportRepliesFromArrivalAddress(t *testing.T) { + if !socket.IPv4ArrivalSupported { + t.Skip("no IPv4 arrival address on this platform") + } + lan := lanIPv4(t) + + udp, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4zero}) + if err != nil { + t.Fatal(err) + } + defer udp.Close() + port := udp.LocalAddr().(*net.UDPAddr).Port + + sock := socket.NewSocket() + m := socket.NewMagicConn(sock, udp) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go m.Serve(ctx) + defer m.Close() + + client, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + server := netip.AddrPortFrom(lan, uint16(port)) + if _, err := client.WriteToUDPAddrPort([]byte("ping"), server); err != nil { + t.Skipf("cannot send from loopback to %s: %v", server, err) + } + + m.SetReadDeadline(time.Now().Add(2 * time.Second)) + buf := make([]byte, 64) + n, from, err := m.ReadFrom(buf) + if err != nil { + t.Skipf("no datagram from loopback to %s arrived: %v", server, err) + } + if string(buf[:n]) != "ping" { + t.Fatalf("received %q", buf[:n]) + } + if _, err := m.WriteTo([]byte("pong"), from); err != nil { + t.Fatalf("WriteTo: %v", err) + } + + client.SetReadDeadline(time.Now().Add(2 * time.Second)) + rn, replyFrom, err := client.ReadFromUDPAddrPort(buf) + if err != nil { + t.Fatalf("client read: %v", err) + } + if string(buf[:rn]) != "pong" { + t.Fatalf("reply %q", buf[:rn]) + } + if replyFrom != server { + t.Fatalf("reply came from %s, want the arrival address %s", replyFrom, server) + } +} diff --git a/internal/socket/ip_pktinfo_unix.go b/internal/socket/ip_pktinfo_unix.go new file mode 100644 index 00000000..43799b49 --- /dev/null +++ b/internal/socket/ip_pktinfo_unix.go @@ -0,0 +1,46 @@ +//go:build unix + +package socket + +import ( + "encoding/binary" + "net/netip" + + "golang.org/x/sys/unix" +) + +// Sizes of struct in_pktinfo and struct in6_pktinfo, laid out the same on +// every platform that has them. +const ( + sizeofInet4Pktinfo = 12 // uint32 ifindex, in_addr spec_dst, in_addr addr + sizeofInet6Pktinfo = 20 // in6_addr addr, uint32 ifindex +) + +// parsePacketInfo extracts the destination address of a received datagram +// from its control messages without allocating. Both families are +// recognized: the level the kernel uses depends on the socket and the +// traffic, not on the bind address alone. +func parsePacketInfo(oob []byte) (localAddr, bool) { + for len(oob) > 0 { + h, data, rest, err := unix.ParseOneSocketControlMessage(oob) + if err != nil { + return localAddr{}, false + } + switch { + case h.Level == unix.IPPROTO_IP && h.Type == sysIPPktinfo && len(data) >= sizeofInet4Pktinfo: + return localAddr{ + addr: netip.AddrFrom4([4]byte(data[8:12])), + ifIndex: int(binary.NativeEndian.Uint32(data[0:4])), + }, true + case h.Level == unix.IPPROTO_IPV6 && h.Type == unix.IPV6_PKTINFO && len(data) >= sizeofInet6Pktinfo: + a := netip.AddrFrom16([16]byte(data[0:16])) + ifIndex := 0 + if a.IsLinkLocalUnicast() { + ifIndex = int(binary.NativeEndian.Uint32(data[16:20])) + } + return localAddr{addr: a.Unmap(), ifIndex: ifIndex}, true + } + oob = rest + } + return localAddr{}, false +} diff --git a/internal/socket/ip_pktinfo_v6_test.go b/internal/socket/ip_pktinfo_v6_test.go new file mode 100644 index 00000000..3e473dc6 --- /dev/null +++ b/internal/socket/ip_pktinfo_v6_test.go @@ -0,0 +1,89 @@ +package socket + +import ( + "context" + "net" + "net/netip" + "testing" + "time" +) + +// pingWildcard binds a MagicConn to the wildcard address of network, sends it +// a datagram from a client bound to clientIP, and waits for it to arrive. +func pingWildcard(t *testing.T, network string, wildcard, clientIP net.IP, server netip.Addr) (*MagicConn, *net.UDPConn) { + t.Helper() + udp, err := net.ListenUDP(network, &net.UDPAddr{IP: wildcard}) + if err != nil { + t.Skipf("listen %s: %v", network, err) + } + t.Cleanup(func() { udp.Close() }) + m := NewMagicConn(NewSocket(), udp) + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + go m.Serve(ctx) + t.Cleanup(func() { m.Close() }) + if !m.transports.ip.pktinfo { + t.Skip("packet info not available on a wildcard socket") + } + client, err := net.ListenUDP(network, &net.UDPAddr{IP: clientIP}) + if err != nil { + t.Skipf("listen client: %v", err) + } + t.Cleanup(func() { client.Close() }) + dst := netip.AddrPortFrom(server, uint16(udp.LocalAddr().(*net.UDPAddr).Port)) + if _, err := client.WriteToUDPAddrPort([]byte("ping"), dst); err != nil { + t.Skipf("send to %s: %v", dst, err) + } + m.SetReadDeadline(time.Now().Add(2 * time.Second)) + buf := make([]byte, 64) + if _, _, err := m.ReadFrom(buf); err != nil { + t.Fatalf("no datagram from %s arrived: %v", client.LocalAddr(), err) + } + return m, client +} + +// roundTrip checks that m recorded client's arrival address as want and that +// a reply to client is delivered. +func roundTrip(t *testing.T, m *MagicConn, client *net.UDPConn, want netip.Addr) { + t.Helper() + clientAddr := canonicalAddrPort(client.LocalAddr().(*net.UDPAddr).AddrPort()) + m.transports.ip.localMu.Lock() + la, ok := m.transports.ip.local[clientAddr] + m.transports.ip.localMu.Unlock() + if !ok || la.addr != want { + t.Fatalf("arrival address for %s = %+v, %v; want %s", clientAddr, la, ok, want) + } + if la.cmsg == nil { + t.Fatalf("no packet-info message built for %s from %s", clientAddr, want) + } + if _, err := m.WriteTo([]byte("pong"), client.LocalAddr()); err != nil { + t.Fatal(err) + } + client.SetReadDeadline(time.Now().Add(2 * time.Second)) + buf := make([]byte, 64) + n, from, err := client.ReadFromUDPAddrPort(buf) + if err != nil { + t.Fatalf("client read: %v", err) + } + if string(buf[:n]) != "pong" { + t.Fatalf("reply %q", buf[:n]) + } + if from.Addr().Unmap() != want { + t.Fatalf("reply came from %s, want %s", from, want) + } +} + +// TestIpTransportRecordsIPv6ArrivalAddress checks that IPv6 packet info is +// parsed and answered from. +func TestIpTransportRecordsIPv6ArrivalAddress(t *testing.T) { + m, client := pingWildcard(t, "udp6", net.IPv6zero, net.IPv6loopback, netip.IPv6Loopback()) + roundTrip(t, m, client, netip.IPv6Loopback()) +} + +// TestIpTransportRecordsDualStackArrivalAddress checks that IPv4 traffic on +// a dual-stack socket is recorded with its IPv4 arrival address, whichever +// level the kernel reports it at. +func TestIpTransportRecordsDualStackArrivalAddress(t *testing.T) { + m, client := pingWildcard(t, "udp", net.IPv6zero, net.IPv4(127, 0, 0, 1), netip.AddrFrom4([4]byte{127, 0, 0, 1})) + roundTrip(t, m, client, netip.AddrFrom4([4]byte{127, 0, 0, 1})) +} diff --git a/internal/socket/transport_gso_linux.go b/internal/socket/transport_gso_linux.go index bddb0f4b..eb445ccf 100644 --- a/internal/socket/transport_gso_linux.go +++ b/internal/socket/transport_gso_linux.go @@ -24,24 +24,41 @@ func (m *MagicConn) WriteMsgUDP(p, oob []byte, addr *net.UDPAddr) (n, oobn int, m.writeMsgSegments(p, addr, segmentSize) return len(p), len(oob), nil } - n, oobn, err = m.udp.WriteMsgUDPAddrPort(p, oob, ap) + // Stack buffer for the packet-info message: no allocation per send. + var cbuf [maxControlSize]byte + msg := m.transports.ip.withPacketInfo(ap, oob, cbuf[:0]) + n, oobn, err = m.udp.WriteMsgUDPAddrPort(p, msg, ap) + if err != nil && len(msg) > len(oob) && !gsoRefused(err, segmentSize) { + // The kernel refused the source: the address it names has gone, or + // the platform will not take it. Drop it and let the kernel choose, + // as IpTransport.send does. Without this the entry outlives the + // address and every later send to this peer is refused the same way. + m.transports.ip.forgetLocal(ap) + n, oobn, err = m.udp.WriteMsgUDPAddrPort(p, oob, ap) + } if err == nil { for range segmentCount(len(p), segmentSize) { m.recordIPSent(ap) } - return n, oobn, nil + return n, len(oob), nil } // EIO on a segmented write is the kernel refusing GSO, not a lost // datagram: the caller disables GSO and resends the batch one datagram at // a time, and those resends are counted on the ordinary path. Counting // here would count them twice. - if segmentSize > 0 && errors.Is(err, unix.EIO) { + if gsoRefused(err, segmentSize) { return n, oobn, err } m.metrics.blackholed.Add(uint64(segmentCount(len(p), segmentSize))) return len(p), len(oob), nil } +// gsoRefused reports whether err is the kernel declining to segment a write +// rather than a fault in the datagram or its source address. +func gsoRefused(err error, segmentSize int) bool { + return segmentSize > 0 && errors.Is(err, unix.EIO) +} + func (m *MagicConn) writeMsgSegments(p []byte, addr *net.UDPAddr, segmentSize int) { if segmentSize <= 0 || segmentSize >= len(p) { m.WriteTo(p, addr)