From 4a8c158b500f9d5844f556e46d4ff234d30e48c9 Mon Sep 17 00:00:00 2001 From: Viktor Liu <17948409+lixmal@users.noreply.github.com> Date: Mon, 28 Sep 2026 16:36:08 +0900 Subject: [PATCH] [client] Resolve the shared socket source address through a connected UDP probe socket (#7633) --- sharedsock/sock_linux.go | 103 ++++++--- sharedsock/src_probe_linux.go | 214 +++++++++++++++++ sharedsock/src_probe_linux_test.go | 218 ++++++++++++++++++ sharedsock/src_probe_privileged_linux_test.go | 92 ++++++++ 4 files changed, 595 insertions(+), 32 deletions(-) create mode 100644 sharedsock/src_probe_linux.go create mode 100644 sharedsock/src_probe_linux_test.go create mode 100644 sharedsock/src_probe_privileged_linux_test.go diff --git a/sharedsock/sock_linux.go b/sharedsock/sock_linux.go index 150e8a722..640f813c9 100644 --- a/sharedsock/sock_linux.go +++ b/sharedsock/sock_linux.go @@ -10,13 +10,13 @@ import ( "context" "fmt" "net" + "net/netip" "time" "github.com/google/gopacket" "github.com/google/gopacket/layers" "github.com/mdlayher/socket" log "github.com/sirupsen/logrus" - "github.com/vishvananda/netlink" "golang.org/x/sync/errgroup" "golang.org/x/sys/unix" @@ -33,6 +33,8 @@ type SharedSocket struct { ctx context.Context conn4 *socket.Conn conn6 *socket.Conn + probe4 *srcProbe + probe6 *srcProbe port int mtu uint16 packetDemux chan rcvdPacket @@ -87,14 +89,26 @@ func Listen(port int, filter BPFFilter, mtu uint16) (_ net.PacketConn, err error return nil, fmt.Errorf("set SO_MARK on ipv4 socket: %w", err) } + if rawSock.probe4, err = newSrcProbe(unix.AF_INET); err != nil { + return nil, err + } + var sockErr error rawSock.conn6, sockErr = socket.Socket(unix.AF_INET6, unix.SOCK_RAW, unix.IPPROTO_UDP, "raw_udp6", nil) if sockErr != nil { - log.Errorf("Failed to create ipv6 raw socket: %v", err) + log.Errorf("Failed to create ipv6 raw socket: %v", sockErr) } else { if err = nbnet.SetSocketMark(rawSock.conn6); err != nil { return nil, fmt.Errorf("set SO_MARK on ipv6 socket: %w", err) } + rawSock.probe6, sockErr = newSrcProbe(unix.AF_INET6) + if sockErr != nil { + log.Errorf("Failed to create ipv6 source probe, continuing without ipv6: %v", sockErr) + if closeErr := rawSock.conn6.Close(); closeErr != nil { + log.Debugf("failed to close ipv6 raw socket: %v", closeErr) + } + rawSock.conn6 = nil + } } ipv4Instructions, ipv6Instructions, err := filter.GetInstructions(uint32(rawSock.port)) @@ -121,23 +135,36 @@ func Listen(port int, filter BPFFilter, mtu uint16) (_ net.PacketConn, err error return rawSock, nil } -// resolveSrc returns the source IP the kernel will pick for a packet sent to -// dst by these raw sockets, mirroring the fwmark the kernel will see on send. -func (s *SharedSocket) resolveSrc(dst net.IP) (net.IP, error) { - opts := &netlink.RouteGetOptions{} - if nbnet.AdvancedRouting() { - opts.Mark = nbnet.ControlPlaneMark +// sockaddr returns the raw send address for dst, carrying the scope of its zone. +func (s *SharedSocket) sockaddr(dst netip.Addr) (unix.Sockaddr, error) { + if dst.Zone() == "" { + return rawSockaddr(dst, 0), nil } - routes, err := netlink.RouteGetWithOptions(dst, opts) + if s.conn6 == nil { + return nil, fmt.Errorf("no raw socket for %s", dst) + } + rc, err := s.conn6.SyscallConn() if err != nil { - return nil, fmt.Errorf("route get %s: %w", dst, err) + return nil, fmt.Errorf("ipv6 raw socket: %w", err) } - for _, r := range routes { - if r.Src != nil { - return r.Src, nil - } + scope, err := zoneIndex(rc, dst.Zone()) + if err != nil { + return nil, err } - return nil, fmt.Errorf("no source IP for %s", dst) + return rawSockaddr(dst, scope), nil +} + +// resolveSrc returns the source IP the kernel will pick for a packet sent to sa +// by these raw sockets, mirroring the fwmark the kernel will see on send. +func (s *SharedSocket) resolveSrc(dst netip.Addr, sa unix.Sockaddr) (netip.Addr, error) { + probe := s.probe4 + if dst.Is6() { + probe = s.probe6 + } + if probe == nil { + return netip.Addr{}, fmt.Errorf("no raw socket for %s", dst) + } + return probe.resolve(sa) } // LocalAddr returns the local address, preferring IPv4 for backward compatibility. @@ -222,6 +249,13 @@ func (s *SharedSocket) Close() error { if s.conn6 != nil { errGrp.Go(s.conn6.Close) } + + if s.probe4 != nil { + errGrp.Go(s.probe4.close) + } + if s.probe6 != nil { + errGrp.Go(s.probe6.close) + } return errGrp.Wait() } @@ -296,14 +330,24 @@ func (s *SharedSocket) WriteTo(buf []byte, rAddr net.Addr) (n int, err error) { DstPort: layers.UDPPort(rUDPAddr.Port), } - src, err := s.resolveSrc(rUDPAddr.IP) - if err != nil { - return 0, fmt.Errorf("resolve source for %s: %w", rUDPAddr.IP, err) + dst := rUDPAddr.AddrPort().Addr().Unmap() + if !dst.IsValid() { + return 0, fmt.Errorf("invalid destination %s", rUDPAddr) } - rSockAddr, conn, nwLayer := s.getWriterObjects(src, rUDPAddr.IP) + rSockAddr, err := s.sockaddr(dst) + if err != nil { + return 0, err + } + + src, err := s.resolveSrc(dst, rSockAddr) + if err != nil { + return 0, fmt.Errorf("resolve source for %s: %w", dst, err) + } + + conn, nwLayer := s.getWriterObjects(src, dst) if conn == nil { - return 0, fmt.Errorf("no raw socket for %s", rUDPAddr.IP) + return 0, fmt.Errorf("no raw socket for %s", dst) } if err := udp.SetNetworkLayerForChecksum(nwLayer); err != nil { @@ -320,28 +364,23 @@ func (s *SharedSocket) WriteTo(buf []byte, rAddr net.Addr) (n int, err error) { } // getWriterObjects returns the specific IP version objects that are used to build a packet and send it using the raw socket -func (s *SharedSocket) getWriterObjects(src, dest net.IP) (sa unix.Sockaddr, conn *socket.Conn, layer gopacket.NetworkLayer) { - if dest.To4() == nil { - sa = &unix.SockaddrInet6{} - copy(sa.(*unix.SockaddrInet6).Addr[:], dest.To16()) +func (s *SharedSocket) getWriterObjects(src, dest netip.Addr) (conn *socket.Conn, layer gopacket.NetworkLayer) { + if dest.Is6() { conn = s.conn6 - layer = &layers.IPv6{ - SrcIP: src, - DstIP: dest, + SrcIP: src.AsSlice(), + DstIP: dest.AsSlice(), } } else { - sa = &unix.SockaddrInet4{} - copy(sa.(*unix.SockaddrInet4).Addr[:], dest.To4()) conn = s.conn4 layer = &layers.IPv4{ Version: 4, TTL: 64, Protocol: layers.IPProtocolUDP, - SrcIP: src, - DstIP: dest, + SrcIP: src.AsSlice(), + DstIP: dest.AsSlice(), } } - return sa, conn, layer + return conn, layer } diff --git a/sharedsock/src_probe_linux.go b/sharedsock/src_probe_linux.go new file mode 100644 index 000000000..5e463b598 --- /dev/null +++ b/sharedsock/src_probe_linux.go @@ -0,0 +1,214 @@ +//go:build linux && !android + +package sharedsock + +import ( + "errors" + "fmt" + "net/netip" + "strconv" + "sync" + "syscall" + "unsafe" + + "github.com/mdlayher/socket" + log "github.com/sirupsen/logrus" + "golang.org/x/sys/unix" + + nbnet "github.com/netbirdio/netbird/client/net" +) + +var errProbeClosed = errors.New("source probe closed") + +// srcProbe finds the source address the kernel picks for a destination by connecting +// a UDP socket that carries the raw sockets' fwmark and reading back its local address. +// Connecting a UDP socket runs the output route lookup without sending anything. +type srcProbe struct { + family int + + mu sync.Mutex + // conn is nil while no socket is open. A failed route lookup keeps the socket, + // any other failure closes it and the next lookup opens a fresh one, so a socket + // in an unknown state is never reused. + conn *socket.Conn + closed bool +} + +// newSrcProbe opens a probe socket for the given address family. +func newSrcProbe(family int) (*srcProbe, error) { + conn, err := openProbeSocket(family) + if err != nil { + return nil, err + } + return &srcProbe{family: family, conn: conn}, nil +} + +// resolve returns the source address the kernel would use for a packet to sa, a +// sockaddr of the probe's family. It is safe for concurrent use. +func (p *srcProbe) resolve(sa unix.Sockaddr) (netip.Addr, error) { + p.mu.Lock() + defer p.mu.Unlock() + + if p.closed { + return netip.Addr{}, errProbeClosed + } + + if p.conn == nil { + conn, err := openProbeSocket(p.family) + if err != nil { + return netip.Addr{}, err + } + p.conn = conn + } + + src, err := p.lookup(sa) + if err != nil { + var rErr *routeError + if !errors.As(err, &rErr) { + if closeErr := p.closeSocket(); closeErr != nil { + log.Debugf("failed to close source probe socket: %v", closeErr) + } + } + return netip.Addr{}, err + } + return src, nil +} + +// close releases the socket. Later lookups fail with errProbeClosed. +func (p *srcProbe) close() error { + p.mu.Lock() + defer p.mu.Unlock() + + p.closed = true + return p.closeSocket() +} + +// closeSocket closes the socket if one is open. Callers must hold p.mu. +func (p *srcProbe) closeSocket() error { + conn := p.conn + p.conn = nil + if conn == nil { + return nil + } + return conn.Close() +} + +// lookup runs one route lookup on the socket. Callers must hold p.mu. +func (p *srcProbe) lookup(sa unix.Sockaddr) (netip.Addr, error) { + rc, err := p.conn.SyscallConn() + if err != nil { + return netip.Addr{}, fmt.Errorf("probe socket: %w", err) + } + + var src netip.Addr + var lookupErr error + if err := rc.Control(func(fd uintptr) { + src, lookupErr = lookupFD(int(fd), sa) + }); err != nil { + return netip.Addr{}, fmt.Errorf("probe socket: %w", err) + } + return src, lookupErr +} + +// routeError is a route lookup the kernel refused. The socket is still usable after it. +type routeError struct { + err error +} + +func (e *routeError) Error() string { + return fmt.Sprintf("route lookup: %v", e.err) +} + +func (e *routeError) Unwrap() error { + return e.err +} + +func lookupFD(fd int, sa unix.Sockaddr) (netip.Addr, error) { + // A connected socket keeps the source address of its first connect and reuses + // it for later route lookups, so dissolve the association first. + if err := disconnect(fd); err != nil { + return netip.Addr{}, fmt.Errorf("disconnect probe socket: %w", err) + } + + if err := unix.Connect(fd, sa); err != nil { + return netip.Addr{}, &routeError{err: err} + } + + local, err := unix.Getsockname(fd) + if err != nil { + return netip.Addr{}, fmt.Errorf("read probe socket address: %w", err) + } + + var src netip.Addr + switch a := local.(type) { + case *unix.SockaddrInet4: + src = netip.AddrFrom4(a.Addr) + case *unix.SockaddrInet6: + src = netip.AddrFrom16(a.Addr) + } + if !src.IsValid() || src.IsUnspecified() { + return netip.Addr{}, &routeError{err: errors.New("no source address")} + } + return src, nil +} + +func openProbeSocket(family int) (*socket.Conn, error) { + conn, err := socket.Socket(family, unix.SOCK_DGRAM, unix.IPPROTO_UDP, "udp_src_probe", nil) + if err != nil { + return nil, fmt.Errorf("create source probe socket: %w", err) + } + + if err := nbnet.SetSocketMark(conn); err != nil { + _ = conn.Close() + return nil, fmt.Errorf("set SO_MARK on source probe socket: %w", err) + } + return conn, nil +} + +// disconnect dissolves a UDP socket's association by connecting to AF_UNSPEC, which +// also clears the source address the kernel pinned on the previous connect. +func disconnect(fd int) error { + sa := unix.RawSockaddr{Family: unix.AF_UNSPEC} + _, _, errno := unix.Syscall(unix.SYS_CONNECT, uintptr(fd), uintptr(unsafe.Pointer(&sa)), unsafe.Sizeof(sa)) + if errno != 0 { + return errno + } + return nil +} + +// rawSockaddr returns the sockaddr for dst with port 0 and the given scope. Port 0 +// matches a raw send, whose route lookup carries no ports. A UDP probe connected +// to it still gets an ephemeral source port before its lookup. +func rawSockaddr(dst netip.Addr, scope uint32) unix.Sockaddr { + if dst.Is4() { + return &unix.SockaddrInet4{Addr: dst.As4()} + } + return &unix.SockaddrInet6{Addr: dst.As16(), ZoneId: scope} +} + +// zoneIndex returns the interface index for an IPv6 zone, which is either an +// interface name or a numeric index. An empty zone is index 0. A name costs one +// SIOCGIFINDEX ioctl on rc, which may be any socket. +func zoneIndex(rc syscall.RawConn, zone string) (uint32, error) { + if zone == "" { + return 0, nil + } + if idx, err := strconv.ParseUint(zone, 10, 32); err == nil { + return uint32(idx), nil + } + + ifr, err := unix.NewIfreq(zone) + if err != nil { + return 0, fmt.Errorf("zone %q: %w", zone, err) + } + var ioctlErr error + if err := rc.Control(func(fd uintptr) { + ioctlErr = unix.IoctlIfreq(int(fd), unix.SIOCGIFINDEX, ifr) + }); err != nil { + return 0, fmt.Errorf("zone %q: %w", zone, err) + } + if ioctlErr != nil { + return 0, fmt.Errorf("resolve zone %q: %w", zone, ioctlErr) + } + return ifr.Uint32(), nil +} diff --git a/sharedsock/src_probe_linux_test.go b/sharedsock/src_probe_linux_test.go new file mode 100644 index 000000000..bdf8e37d8 --- /dev/null +++ b/sharedsock/src_probe_linux_test.go @@ -0,0 +1,218 @@ +//go:build linux && !android + +package sharedsock + +import ( + "net" + "net/netip" + "strconv" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/vishvananda/netlink" + "golang.org/x/sys/unix" +) + +// routeGetSrc is the kernel's answer through netlink, used as the reference. +func routeGetSrc(t *testing.T, dst netip.Addr) (netip.Addr, bool) { + t.Helper() + routes, err := netlink.RouteGet(net.IP(dst.AsSlice())) + if err != nil { + return netip.Addr{}, false + } + for _, r := range routes { + if src, ok := netip.AddrFromSlice(r.Src); ok { + return src.Unmap(), true + } + } + return netip.Addr{}, false +} + +func newTestProbe(t *testing.T, family int) *srcProbe { + t.Helper() + p, err := newSrcProbe(family) + require.NoError(t, err) + t.Cleanup(func() { _ = p.close() }) + return p +} + +// A reused UDP socket keeps the source address of its first connect. Alternating +// between a loopback and an off-host destination catches a probe that forgets to +// disconnect: the second lookup would report 127.0.0.1 for the off-host address. +func TestSrcProbe_AlternatingDestinationsMatchRouteGet(t *testing.T) { + loopback := netip.MustParseAddr("127.0.0.1") + remote := netip.MustParseAddr("192.0.2.1") + + remoteSrc, ok := routeGetSrc(t, remote) + if !ok { + t.Skip("no route to an off-host IPv4 destination") + } + require.NotEqual(t, loopback, remoteSrc, "off-host destination must not route via loopback") + + p := newTestProbe(t, unix.AF_INET) + for i := 0; i < 3; i++ { + src, err := p.resolve(rawSockaddr(loopback, 0)) + require.NoError(t, err) + assert.Equal(t, loopback, src, "source for loopback, round %d", i) + + src, err = p.resolve(rawSockaddr(remote, 0)) + require.NoError(t, err) + assert.Equal(t, remoteSrc, src, "source for %s, round %d", remote, i) + } +} + +func TestSrcProbe_IPv6(t *testing.T) { + loopback := netip.MustParseAddr("::1") + if _, ok := routeGetSrc(t, loopback); !ok { + t.Skip("no IPv6 loopback") + } + + p := newTestProbe(t, unix.AF_INET6) + src, err := p.resolve(rawSockaddr(loopback, 0)) + require.NoError(t, err) + assert.Equal(t, loopback, src, "source for ::1") + + remote := netip.MustParseAddr("2001:db8::1") + remoteSrc, ok := routeGetSrc(t, remote) + if !ok { + t.Skipf("no route to %s", remote) + } + src, err = p.resolve(rawSockaddr(remote, 0)) + require.NoError(t, err) + assert.Equal(t, remoteSrc, src, "source for %s", remote) +} + +// A route the kernel refuses is an answer, not a broken socket, so the probe keeps +// its socket and the next lookup reuses it. +func TestSrcProbe_RouteErrorKeepsSocket(t *testing.T) { + p := newTestProbe(t, unix.AF_INET) + conn := p.conn + + // Connecting to the limited broadcast address without SO_BROADCAST fails. + _, err := p.resolve(rawSockaddr(netip.MustParseAddr("255.255.255.255"), 0)) + var rErr *routeError + require.ErrorAs(t, err, &rErr, "lookup for the broadcast address should fail as a route error") + assert.Same(t, conn, p.conn, "route error should keep the probe socket") + + src, err := p.resolve(rawSockaddr(netip.MustParseAddr("127.0.0.1"), 0)) + require.NoError(t, err) + assert.Equal(t, netip.MustParseAddr("127.0.0.1"), src, "source after the route error") + assert.Same(t, conn, p.conn, "lookup after a route error should reuse the socket") +} + +// A socket that stops working is dropped, and the next lookup opens a fresh one. +func TestSrcProbe_ReopensAfterSocketError(t *testing.T) { + p := newTestProbe(t, unix.AF_INET) + require.NoError(t, p.conn.Close()) + + _, err := p.resolve(rawSockaddr(netip.MustParseAddr("127.0.0.1"), 0)) + require.Error(t, err, "lookup on a closed socket should fail") + assert.Nil(t, p.conn, "socket error should drop the probe socket") + + src, err := p.resolve(rawSockaddr(netip.MustParseAddr("127.0.0.1"), 0)) + require.NoError(t, err) + assert.Equal(t, netip.MustParseAddr("127.0.0.1"), src, "source after reopening") + assert.NotNil(t, p.conn, "probe socket should be open again") +} + +// A link-local destination is only routable with its scope. The probe must pass the +// scope to the kernel and get the interface's own link-local address back. +func TestSrcProbe_LinkLocalWithZone(t *testing.T) { + iface, want := linkLocalInterface(t) + dst := netip.MustParseAddr("fe80::1") + + p := newTestProbe(t, unix.AF_INET6) + rc, err := p.conn.SyscallConn() + require.NoError(t, err) + + for _, zone := range []string{iface.Name, strconv.Itoa(iface.Index)} { + scope, err := zoneIndex(rc, zone) + require.NoError(t, err, "zone %q", zone) + assert.Equal(t, uint32(iface.Index), scope, "index for zone %q", zone) + + src, err := p.resolve(rawSockaddr(dst.WithZone(zone), scope)) + require.NoError(t, err, "zone %q", zone) + assert.Equal(t, want, src, "source for %s%%%s", dst, zone) + } + + _, err = p.resolve(rawSockaddr(dst, 0)) + var rErr *routeError + assert.ErrorAs(t, err, &rErr, "link-local destination without a scope should fail as a route error") +} + +func TestZoneIndex(t *testing.T) { + p := newTestProbe(t, unix.AF_INET6) + rc, err := p.conn.SyscallConn() + require.NoError(t, err) + + scope, err := zoneIndex(rc, "") + require.NoError(t, err) + assert.Zero(t, scope, "empty zone should be index 0") + + _, err = zoneIndex(rc, "nb-no-such-if0") + assert.Error(t, err, "unknown interface name should fail") +} + +// linkLocalInterface returns an up interface and its IPv6 link-local address. +func linkLocalInterface(t *testing.T) (net.Interface, netip.Addr) { + t.Helper() + ifaces, err := net.Interfaces() + require.NoError(t, err) + for _, iface := range ifaces { + if iface.Flags&net.FlagUp == 0 || iface.Flags&net.FlagLoopback != 0 { + continue + } + addrs, err := iface.Addrs() + if err != nil { + continue + } + for _, a := range addrs { + prefix, err := netip.ParsePrefix(a.String()) + if err == nil && prefix.Addr().IsLinkLocalUnicast() { + return iface, prefix.Addr() + } + } + } + t.Skip("no interface with an IPv6 link-local address") + return net.Interface{}, netip.Addr{} +} + +func TestSrcProbe_ClosedRejectsLookups(t *testing.T) { + p, err := newSrcProbe(unix.AF_INET) + require.NoError(t, err) + require.NoError(t, p.close()) + require.NoError(t, p.close(), "close must be idempotent") + + _, err = p.resolve(rawSockaddr(netip.MustParseAddr("127.0.0.1"), 0)) + assert.ErrorIs(t, err, errProbeClosed) + assert.Nil(t, p.conn, "closed probe must not reopen") +} + +func BenchmarkSrcProbe(b *testing.B) { + p, err := newSrcProbe(unix.AF_INET) + require.NoError(b, err) + defer p.close() + + dst := netip.MustParseAddr("192.0.2.1") + if _, err := p.resolve(rawSockaddr(dst, 0)); err != nil { + b.Skipf("no route to %s: %v", dst, err) + } + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + if _, err := p.resolve(rawSockaddr(dst, 0)); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkSrcRouteGet(b *testing.B) { + dst := net.ParseIP("192.0.2.1") + b.ReportAllocs() + for i := 0; i < b.N; i++ { + if _, err := netlink.RouteGetWithOptions(dst, &netlink.RouteGetOptions{}); err != nil { + b.Skipf("no route to %s: %v", dst, err) + } + } +} diff --git a/sharedsock/src_probe_privileged_linux_test.go b/sharedsock/src_probe_privileged_linux_test.go new file mode 100644 index 000000000..fed1d3985 --- /dev/null +++ b/sharedsock/src_probe_privileged_linux_test.go @@ -0,0 +1,92 @@ +//go:build privileged + +package sharedsock + +import ( + "net" + "net/netip" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/vishvananda/netlink" + "golang.org/x/sys/unix" + + nbnet "github.com/netbirdio/netbird/client/net" +) + +// The probe must pick the source the raw sockets get on send, which carry the +// control-plane fwmark. The test installs the same shape of policy rule the client +// uses: unmarked traffic to 192.0.2.0/24 is diverted into a table that routes it +// via loopback, so an unmarked lookup reports 127.0.0.1 and a marked one does not. +func TestSrcProbe_HonoursControlPlaneMark(t *testing.T) { + nbnet.Init() + if !nbnet.AdvancedRouting() { + t.Skip("advanced routing not supported") + } + + const ( + table = 4242 + // Below the client's own rules, so a default route in the netbird table cannot win. + priority = 90 + ) + dst := netip.MustParseAddr("192.0.2.1") + + marked, err := netlink.RouteGetWithOptions(net.IP(dst.AsSlice()), &netlink.RouteGetOptions{Mark: nbnet.ControlPlaneMark}) + if err != nil { + t.Skipf("no route to %s: %v", dst, err) + } + if len(marked) == 0 || marked[0].Src == nil { + t.Skipf("marked route to %s has no source address", dst) + } + markedSrc, ok := netip.AddrFromSlice(marked[0].Src) + require.True(t, ok, "parse marked source") + markedSrc = markedSrc.Unmap() + if markedSrc == netip.MustParseAddr("127.0.0.1") { + t.Skipf("marked route to %s already uses loopback, no contrast to test", dst) + } + + rules, err := netlink.RuleList(unix.AF_INET) + require.NoError(t, err) + for _, r := range rules { + if r.Priority == priority || r.Table == table { + t.Skipf("rule priority %d or table %d already in use", priority, table) + } + } + + lo, err := netlink.LinkByName("lo") + require.NoError(t, err) + + route := &netlink.Route{ + Dst: &net.IPNet{IP: net.IPv4(192, 0, 2, 0), Mask: net.CIDRMask(24, 32)}, + LinkIndex: lo.Attrs().Index, + // 127.0.0.1 is host-scoped, so a link-scoped route only picks it when told to. + Src: net.IPv4(127, 0, 0, 1), + Table: table, + Scope: netlink.SCOPE_LINK, + } + require.NoError(t, netlink.RouteAdd(route)) + t.Cleanup(func() { _ = netlink.RouteDel(route) }) + + rule := netlink.NewRule() + rule.Family = unix.AF_INET + rule.Priority = priority + rule.Table = table + rule.Mark = nbnet.ControlPlaneMark + rule.Invert = true + require.NoError(t, netlink.RuleAdd(rule)) + t.Cleanup(func() { _ = netlink.RuleDel(rule) }) + + unmarked, err := netlink.RouteGet(net.IP(dst.AsSlice())) + require.NoError(t, err) + require.NotEmpty(t, unmarked) + require.True(t, unmarked[0].Src.Equal(net.IPv4(127, 0, 0, 1)), "unmarked lookup should be diverted to loopback, got %s", unmarked[0].Src) + + p, err := newSrcProbe(unix.AF_INET) + require.NoError(t, err) + t.Cleanup(func() { _ = p.close() }) + + src, err := p.resolve(rawSockaddr(dst, 0)) + require.NoError(t, err) + assert.Equal(t, markedSrc, src, "probe should resolve the source of the marked lookup") +}