[client] Resolve the shared socket source address through a connected UDP probe socket (#7633)

This commit is contained in:
Viktor Liu
2026-09-28 09:36:08 +02:00
committed by GitHub
parent d5488db146
commit 4a8c158b50
4 changed files with 595 additions and 32 deletions
+71 -32
View File
@@ -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
}