mirror of
https://github.com/netbirdio/netbird.git
synced 2026-09-28 09:39:05 +02:00
[client] Resolve the shared socket source address through a connected UDP probe socket (#7633)
This commit is contained in:
+71
-32
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
Reference in New Issue
Block a user