mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-05 04:59:06 +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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user