Merge remote-tracking branch 'origin/master' into multiport

This commit is contained in:
Wade Simmons
2026-07-23 10:57:32 -04:00
47 changed files with 2706 additions and 693 deletions
+61
View File
@@ -0,0 +1,61 @@
package udp
import (
"context"
"log/slog"
"github.com/slackhq/nebula/config"
)
// NetworkChangeMonitor rebinds the udp listener when the local network moves out from under it.
//
// Detection lives here in the udp package, next to the socket it concerns and the platform matrix that already knows
// which sockets go stale. What to do about a change — updating the lighthouse, requerying tunnels — is not the udp
// package's business, so Start takes the reaction as a plain function. Passing it at Start rather than holding it
// keeps this package from referencing whatever owns the rebind.
//
// On platforms whose sockets do not go stale, watchNetworkChanges hands back a nil channel and Start returns.
type NetworkChangeMonitor struct {
l *slog.Logger
ctx context.Context
enabled bool
}
// NewNetworkChangeMonitor builds a monitor for local network changes. The returned monitor is always usable: Start
// is safe to call unconditionally, it no-ops when disabled or on a platform that does not need it.
func NewNetworkChangeMonitor(ctx context.Context, l *slog.Logger, c *config.C) *NetworkChangeMonitor {
return &NetworkChangeMonitor{
l: l,
ctx: ctx,
enabled: c.GetBool("listen.rebind_on_network_change", true),
}
}
// Start watches for network changes until the context is cancelled, calling rebind once per settled change. It
// blocks, so callers run it in a goroutine, and it no-ops when disabled, unsupported, or with nothing to rebind.
func (m *NetworkChangeMonitor) Start(rebind func()) {
if !m.enabled || rebind == nil || m.ctx.Err() != nil {
return
}
changes, err := watchNetworkChanges(m.ctx, m.l)
if err != nil {
// Not fatal. Everything else still works, we just won't notice a network change on our own.
m.l.Error("Failed to watch for network changes, will not rebind the udp listener when the network moves",
"error", err,
)
return
}
if changes == nil {
// This platform's sockets don't go stale, so there is nothing to watch for.
return
}
m.l.Info("Watching for network changes to rebind the udp listener")
for range changes {
m.l.Info("Local network changed, rebinding the udp listener")
rebind()
}
}
+164
View File
@@ -0,0 +1,164 @@
//go:build darwin && !ios && !e2e_testing
// +build darwin,!ios,!e2e_testing
package udp
import (
"context"
"encoding/binary"
"errors"
"log/slog"
"os"
"time"
"golang.org/x/sys/unix"
)
const (
// netChangeSettleWindow is how long we keep swallowing routing messages after the first interesting one. A
// single network change is never a single message, it is a burst: the link drops, addresses go away, new ones
// arrive, routes get rewritten. Reporting part way through that just means reporting again.
netChangeSettleWindow = time.Second
// netChangeReadBuffer is sized well past any rt_msghdr plus its addresses. A short read would be discarded by
// the kernel, so being generous here is how we avoid missing a message.
netChangeReadBuffer = 4096
)
// watchNetworkChanges reports when the local network moves out from under us, so the listener can be rebound.
//
// Darwin scopes a udp socket to whatever interface it came up on. Move between networks and we keep sending out an
// interface that no longer has a route, which surfaces as an instant "no route to host" with no packet ever leaving
// the box. Rebind clears that, but only if something notices the change and calls it. iOS has always been told by
// the host app off NWPathMonitor. This is the equivalent for everything else that runs on darwin.
//
// The returned channel is buffered and coalescing: a send is dropped if one is already pending, since both mean the
// same thing to a reader. It is closed when ctx is cancelled or the routing socket fails, so a caller can simply
// range over it. Platforms whose sockets do not need rebinding return a nil channel and no error.
func watchNetworkChanges(ctx context.Context, l *slog.Logger) (<-chan struct{}, error) {
sock, err := openRouteSocket()
if err != nil {
return nil, err
}
changes := make(chan struct{}, 1)
go func() {
defer close(changes)
defer func() { _ = sock.Close() }()
// Closing the socket is what unblocks the read in watchRouteSocket, so this turns cancellation into a
// close. It is scoped to this call so it cannot outlive the watch it belongs to.
done := make(chan struct{})
defer close(done)
go func() {
select {
case <-ctx.Done():
_ = sock.Close()
case <-done:
}
}()
watchRouteSocket(l, sock, changes)
}()
return changes, nil
}
// watchRouteSocket blocks reading the routing socket, reporting once per settled burst of changes. It returns when
// the socket is closed, which is how cancellation gets us out of here.
func watchRouteSocket(l *slog.Logger, sock *os.File, changes chan<- struct{}) {
buf := make([]byte, netChangeReadBuffer)
for {
n, err := sock.Read(buf)
if err != nil {
logRouteSocketError(l, err)
return
}
if !isNetworkChange(buf[:n]) {
continue
}
// Swallow the rest of the burst. The deadline is absolute and not extended by what arrives, so this always
// ends after the settle window no matter how chatty the socket is. Changes that land after the window
// simply produce another report, which is the correct outcome anyway.
deadline := time.Now().Add(netChangeSettleWindow)
for {
if err = sock.SetReadDeadline(deadline); err != nil {
logRouteSocketError(l, err)
return
}
if _, err = sock.Read(buf); err != nil {
if os.IsTimeout(err) {
break
}
logRouteSocketError(l, err)
return
}
}
if err = sock.SetReadDeadline(time.Time{}); err != nil {
logRouteSocketError(l, err)
return
}
select {
case changes <- struct{}{}:
default:
// One already pending, and a second "the network moved" tells the reader nothing new.
}
}
}
// logRouteSocketError reports a routing socket failure unless it is just us shutting the socket down.
func logRouteSocketError(l *slog.Logger, err error) {
if errors.Is(err, os.ErrClosed) {
return
}
l.Error("Error reading the routing socket, will no longer notice local network changes", "error", err)
}
// openRouteSocket returns the routing socket as a non blocking os.File. Going through os.File puts reads on the go
// poller, which buys us both a working read deadline and a Close that unblocks a read in progress.
func openRouteSocket() (*os.File, error) {
fd, err := unix.Socket(unix.AF_ROUTE, unix.SOCK_RAW, unix.AF_UNSPEC)
if err != nil {
return nil, err
}
if err = unix.SetNonblock(fd, true); err != nil {
_ = unix.Close(fd)
return nil, err
}
return os.NewFile(uintptr(fd), "route"), nil
}
// isNetworkChange reports whether a routing message means our local addressing may have moved out from under us.
//
// We read the header instead of parsing the message because the type is the only part we need, and a full parse can
// fail on shapes we don't care about, which would turn "a message I can't parse" into "a change I missed".
// rt_msghdr, if_msghdr and ifa_msghdr all begin with the same three fields, so this is the same for every type.
func isNetworkChange(msg []byte) bool {
if len(msg) < 4 {
return false
}
// u_short msglen, u_char version, u_char type
if int(binary.NativeEndian.Uint16(msg[0:2])) > len(msg) || msg[2] != unix.RTM_VERSION {
return false
}
switch msg[3] {
case unix.RTM_NEWADDR, unix.RTM_DELADDR, unix.RTM_IFINFO:
// An address arrived or left, or a link changed state. Anything else on this socket is either a route
// churning underneath us, which a rebind doesn't help with, or unrelated traffic.
return true
default:
return false
}
}
+244
View File
@@ -0,0 +1,244 @@
//go:build darwin && !ios && !e2e_testing
// +build darwin,!ios,!e2e_testing
package udp
import (
"context"
"encoding/binary"
"os"
"testing"
"time"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/test"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/goleak"
"golang.org/x/sys/unix"
)
// routeMsg builds the first four bytes of a routing message, which is all isNetworkChange reads.
func routeMsg(msgType uint8, extra int) []byte {
msg := make([]byte, 4+extra)
binary.NativeEndian.PutUint16(msg[0:2], uint16(len(msg)))
msg[2] = unix.RTM_VERSION
msg[3] = msgType
return msg
}
func TestIsNetworkChange(t *testing.T) {
// The three that mean our addressing may have moved
assert.True(t, isNetworkChange(routeMsg(unix.RTM_NEWADDR, 0)))
assert.True(t, isNetworkChange(routeMsg(unix.RTM_DELADDR, 0)))
assert.True(t, isNetworkChange(routeMsg(unix.RTM_IFINFO, 0)))
// Route churn is not something a rebind helps with
assert.False(t, isNetworkChange(routeMsg(unix.RTM_ADD, 0)))
assert.False(t, isNetworkChange(routeMsg(unix.RTM_DELETE, 0)))
assert.False(t, isNetworkChange(routeMsg(unix.RTM_GET, 0)))
// Garbage must not be mistaken for a change
assert.False(t, isNetworkChange(nil), "empty")
assert.False(t, isNetworkChange([]byte{0, 0, 0}), "short header")
wrongVersion := routeMsg(unix.RTM_NEWADDR, 0)
wrongVersion[2] = unix.RTM_VERSION + 1
assert.False(t, isNetworkChange(wrongVersion), "wrong rtm_version")
lying := routeMsg(unix.RTM_NEWADDR, 0)
binary.NativeEndian.PutUint16(lying[0:2], 512)
assert.False(t, isNetworkChange(lying), "msglen longer than what we read")
}
// socketPair returns a connected pair of datagram sockets, the first wrapped the same way the routing socket is. It
// stands in for the kernel so the watch loop can be driven with synthetic messages.
func socketPair(t *testing.T) (*os.File, int) {
t.Helper()
fds, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_DGRAM, 0)
require.NoError(t, err)
require.NoError(t, unix.SetNonblock(fds[0], true))
f := os.NewFile(uintptr(fds[0]), "route")
t.Cleanup(func() {
_ = f.Close()
_ = unix.Close(fds[1])
})
return f, fds[1]
}
func TestWatchRouteSocketCoalescesABurst(t *testing.T) {
sock, kernel := socketPair(t)
changes := make(chan struct{}, 1)
done := make(chan struct{})
go func() {
watchRouteSocket(test.NewLogger(), sock, changes)
close(done)
}()
// One network change is a burst of messages. All of these land inside the settle window, so they must produce
// exactly one report rather than one apiece.
for range 5 {
_, err := unix.Write(kernel, routeMsg(unix.RTM_NEWADDR, 8))
require.NoError(t, err)
}
// Uninteresting messages in the middle of a burst must not add a report of their own either.
_, err := unix.Write(kernel, routeMsg(unix.RTM_ADD, 8))
require.NoError(t, err)
select {
case <-changes:
case <-time.After(netChangeSettleWindow * 4):
t.Fatal("a burst should have reported a change")
}
// Nothing more from that burst
select {
case <-changes:
t.Fatal("a burst should report exactly once")
case <-time.After(netChangeSettleWindow):
}
// A change after the window has closed is a separate event and gets its own report.
_, err = unix.Write(kernel, routeMsg(unix.RTM_IFINFO, 8))
require.NoError(t, err)
select {
case <-changes:
case <-time.After(netChangeSettleWindow * 4):
t.Fatal("a later change should report again")
}
// Closing the socket is how the real thing shuts down
require.NoError(t, sock.Close())
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("watchRouteSocket did not return after the socket was closed")
}
}
func TestWatchRouteSocketIgnoresUninterestingMessages(t *testing.T) {
sock, kernel := socketPair(t)
changes := make(chan struct{}, 1)
done := make(chan struct{})
go func() {
watchRouteSocket(test.NewLogger(), sock, changes)
close(done)
}()
for _, msgType := range []uint8{unix.RTM_ADD, unix.RTM_DELETE, unix.RTM_GET, unix.RTM_MISS} {
_, err := unix.Write(kernel, routeMsg(msgType, 8))
require.NoError(t, err)
}
select {
case <-changes:
t.Fatal("route churn alone must not report a change")
case <-time.After(netChangeSettleWindow * 2):
}
require.NoError(t, sock.Close())
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("watchRouteSocket did not return after the socket was closed")
}
}
// TestWatchRouteSocketDropsRatherThanBlocks covers the coalescing send. A reader that is busy rebinding must not
// wedge the watcher, and a second pending "the network moved" tells it nothing new anyway.
func TestWatchRouteSocketDropsRatherThanBlocks(t *testing.T) {
sock, kernel := socketPair(t)
changes := make(chan struct{}, 1)
done := make(chan struct{})
go func() {
watchRouteSocket(test.NewLogger(), sock, changes)
close(done)
}()
// Nobody is reading changes, so after the first report the buffer is full for the rest of this test
for range 3 {
_, err := unix.Write(kernel, routeMsg(unix.RTM_NEWADDR, 8))
require.NoError(t, err)
time.Sleep(netChangeSettleWindow + time.Millisecond*250)
}
// The watcher must still be alive and responsive to a close
require.NoError(t, sock.Close())
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("watchRouteSocket wedged on a full channel")
}
assert.Len(t, changes, 1, "the pending report should have coalesced, not queued")
}
// TestWatchNetworkChangesStopsWithContext covers the detection path against a real routing socket, including that
// cancelling the context closes the channel so a ranging caller falls out of its loop.
func TestWatchNetworkChangesStopsWithContext(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
changes, err := watchNetworkChanges(ctx, test.NewLogger())
require.NoError(t, err)
require.NotNil(t, changes, "darwin should support watching")
drained := make(chan struct{})
go func() {
for range changes {
}
close(drained)
}()
cancel()
select {
case <-drained:
case <-time.After(time.Second * 5):
t.Fatal("cancelling the context should close the changes channel")
}
}
// TestNetworkChangeMonitorStopsWithContext drives the whole monitor against a real routing socket: Start must block
// watching, and cancelling the context (which is all Control does on shutdown, it never stops the monitor directly)
// must return it and clean up the watch goroutines.
func TestNetworkChangeMonitorStopsWithContext(t *testing.T) {
// IgnoreCurrent because other tests in this package leave readers running; we only care about what this test
// leaks itself.
defer goleak.VerifyNone(t, goleak.IgnoreCurrent())
ctx, cancel := context.WithCancel(context.Background())
l := test.NewLogger()
c := config.NewC(l)
require.NoError(t, c.LoadString("listen:\n rebind_on_network_change: true\n"))
m := NewNetworkChangeMonitor(ctx, l, c)
done := make(chan struct{})
go func() {
m.Start(func() {})
close(done)
}()
// Start should be sitting on the routing socket, not have fallen out. If it returned early it either failed to
// watch or no-op'd, both of which we want to catch.
select {
case <-done:
t.Fatal("Start returned instead of watching")
case <-time.After(time.Millisecond * 250):
}
cancel()
select {
case <-done:
case <-time.After(time.Second * 5):
t.Fatal("Start did not return after the context was cancelled")
}
// Starting again after the context is dead must not open anything.
m.Start(func() {})
}
+22
View File
@@ -0,0 +1,22 @@
//go:build !darwin || ios || e2e_testing
// +build !darwin ios e2e_testing
package udp
import (
"context"
"log/slog"
)
// watchNetworkChanges is a no-op outside of darwin.
//
// Darwin is the platform that scopes a udp socket to the interface it came up on, so it is the platform whose socket
// goes stale when the local network changes. Everywhere else Rebind has nothing to do, so there is nothing to watch
// for. iOS is excluded on purpose even though it is darwin: the host app already drives the rebind off NWPathMonitor,
// and two things racing to rebind the same socket is worse than one.
//
// A nil channel means "not supported here", which callers must treat as "do not start a watcher" rather than
// selecting on it, since a receive from a nil channel blocks forever.
func watchNetworkChanges(_ context.Context, _ *slog.Logger) (<-chan struct{}, error) {
return nil, nil
}
+39
View File
@@ -0,0 +1,39 @@
package udp
import (
"context"
"testing"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/test"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func newMonitor(t *testing.T, ctx context.Context, cfg string) *NetworkChangeMonitor {
t.Helper()
l := test.NewLogger()
c := config.NewC(l)
require.NoError(t, c.LoadString(cfg))
return NewNetworkChangeMonitor(ctx, l, c)
}
func TestNetworkChangeMonitorDefaultsOn(t *testing.T) {
// Says nothing about rebinding, so this covers the default.
m := newMonitor(t, context.Background(), "listen:\n host: 0.0.0.0\n")
assert.True(t, m.enabled, "should default to on")
}
func TestNetworkChangeMonitorDisabledIsANoOp(t *testing.T) {
m := newMonitor(t, context.Background(), "listen:\n rebind_on_network_change: false\n")
require.False(t, m.enabled)
// Must return without opening a socket. If it watched anything this would block.
m.Start(func() {})
}
func TestNetworkChangeMonitorNilRebindIsANoOp(t *testing.T) {
// Nothing to rebind, so there is no point watching, on any platform.
m := newMonitor(t, context.Background(), "listen:\n rebind_on_network_change: true\n")
m.Start(nil)
}
+4 -5
View File
@@ -187,6 +187,9 @@ func (u *StdConn) SupportsMultipleReaders() bool {
return false
}
// Rebind clears the interface the kernel scoped this socket to, so that sends are routed against the current
// routing table instead of the interface we happened to be on when the socket was created. Darwin pins sockets
// this way on its own, which is what strands us after the underlying network changes.
func (u *StdConn) Rebind() error {
var err error
if u.isV4 {
@@ -195,9 +198,5 @@ func (u *StdConn) Rebind() error {
err = syscall.SetsockoptInt(int(u.sysFd), syscall.IPPROTO_IPV6, syscall.IPV6_BOUND_IF, 0)
}
if err != nil {
u.l.Error("Failed to rebind udp socket", "error", err)
}
return nil
return err
}
+168 -167
View File
@@ -4,12 +4,13 @@
package udp
import (
"context"
"encoding/binary"
"errors"
"fmt"
"log/slog"
"net"
"net/netip"
"sync/atomic"
"syscall"
"unsafe"
@@ -19,58 +20,51 @@ import (
)
type StdConn struct {
udpConn *net.UDPConn
rawConn syscall.RawConn
isV4 bool
l *slog.Logger
batch int
}
func setReusePort(network, address string, c syscall.RawConn) error {
var opErr error
err := c.Control(func(fd uintptr) {
opErr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_REUSEPORT, 1)
//CloseOnExec already set by the runtime
})
if err != nil {
return err
}
return opErr
sysFd int
closed atomic.Bool
isV4 bool
l *slog.Logger
batch int
}
func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) {
listen := netip.AddrPortFrom(ip, uint16(port))
lc := net.ListenConfig{}
af := unix.AF_INET6
if ip.Is4() {
af = unix.AF_INET
}
syscall.ForkLock.RLock()
fd, err := unix.Socket(af, unix.SOCK_DGRAM, unix.IPPROTO_UDP)
if err == nil {
unix.CloseOnExec(fd)
}
syscall.ForkLock.RUnlock()
if err != nil {
return nil, fmt.Errorf("unable to open socket: %w", err)
}
if multi {
lc.Control = setReusePort
}
//this context is only used during the bind operation, you can't cancel it to kill the socket
pc, err := lc.ListenPacket(context.Background(), "udp", listen.String())
if err != nil {
return nil, fmt.Errorf("unable to open socket: %s", err)
}
udpConn := pc.(*net.UDPConn)
rawConn, err := udpConn.SyscallConn()
if err != nil {
_ = udpConn.Close()
return nil, err
}
//gotta find out if we got an AF_INET6 socket or not:
out := &StdConn{
udpConn: udpConn,
rawConn: rawConn,
l: l,
batch: batch,
if err = unix.SetsockoptInt(fd, unix.SOL_SOCKET, unix.SO_REUSEPORT, 1); err != nil {
_ = unix.Close(fd)
return nil, fmt.Errorf("unable to set SO_REUSEPORT: %w", err)
}
}
af, err := out.getSockOptInt(unix.SO_DOMAIN)
if err != nil {
_ = out.Close()
return nil, err
var sa unix.Sockaddr
if ip.Is4() {
sa4 := &unix.SockaddrInet4{Port: port}
sa4.Addr = ip.As4()
sa = sa4
} else {
sa6 := &unix.SockaddrInet6{Port: port}
sa6.Addr = ip.As16()
sa = sa6
}
if err = unix.Bind(fd, sa); err != nil {
_ = unix.Close(fd)
return nil, fmt.Errorf("unable to bind to socket: %w", err)
}
out.isV4 = af == unix.AF_INET
return out, nil
return &StdConn{sysFd: fd, isV4: ip.Is4(), l: l, batch: batch}, nil
}
func (u *StdConn) SupportsMultipleReaders() bool {
@@ -81,134 +75,111 @@ func (u *StdConn) Rebind() error {
return nil
}
func (u *StdConn) getSockOptInt(opt int) (int, error) {
if u.rawConn == nil {
return 0, fmt.Errorf("no UDP connection")
}
var out int
var opErr error
err := u.rawConn.Control(func(fd uintptr) {
out, opErr = unix.GetsockoptInt(int(fd), unix.SOL_SOCKET, opt)
})
if err != nil {
return 0, err
}
return out, opErr
}
func (u *StdConn) setSockOptInt(opt int, n int) error {
if u.rawConn == nil {
return fmt.Errorf("no UDP connection")
}
var opErr error
err := u.rawConn.Control(func(fd uintptr) {
opErr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, opt, n)
})
if err != nil {
return err
}
return opErr
}
func (u *StdConn) SetRecvBuffer(n int) error {
return u.setSockOptInt(unix.SO_RCVBUFFORCE, n)
return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_RCVBUFFORCE, n)
}
func (u *StdConn) SetSendBuffer(n int) error {
return u.setSockOptInt(unix.SO_SNDBUFFORCE, n)
return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_SNDBUFFORCE, n)
}
func (u *StdConn) SetSoMark(mark int) error {
return u.setSockOptInt(unix.SO_MARK, mark)
return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_MARK, mark)
}
func (u *StdConn) GetRecvBuffer() (int, error) {
return u.getSockOptInt(unix.SO_RCVBUF)
return unix.GetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_RCVBUF)
}
func (u *StdConn) GetSendBuffer() (int, error) {
return u.getSockOptInt(unix.SO_SNDBUF)
return unix.GetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_SNDBUF)
}
func (u *StdConn) GetSoMark() (int, error) {
return u.getSockOptInt(unix.SO_MARK)
return unix.GetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_MARK)
}
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
a := u.udpConn.LocalAddr()
switch v := a.(type) {
case *net.UDPAddr:
addr, ok := netip.AddrFromSlice(v.IP)
if !ok {
return netip.AddrPort{}, fmt.Errorf("LocalAddr returned invalid IP address: %s", v.IP)
}
return netip.AddrPortFrom(addr, uint16(v.Port)), nil
sa, err := unix.Getsockname(u.sysFd)
if err != nil {
return netip.AddrPort{}, err
}
switch sa := sa.(type) {
case *unix.SockaddrInet4:
return netip.AddrPortFrom(netip.AddrFrom4(sa.Addr), uint16(sa.Port)), nil
case *unix.SockaddrInet6:
return netip.AddrPortFrom(netip.AddrFrom16(sa.Addr), uint16(sa.Port)), nil
default:
return netip.AddrPort{}, fmt.Errorf("LocalAddr returned: %#v", a)
return netip.AddrPort{}, fmt.Errorf("unsupported sock type: %T", sa)
}
}
func recvmmsg(fd uintptr, msgs []rawMessage) (int, bool, error) {
var errno syscall.Errno
n, _, errno := unix.Syscall6(
// recvmmsg does one blocking recvmmsg (MSG_WAITFORONE), reading up to len(msgs) datagrams
func (u *StdConn) recvmmsg(msgs []rawMessage) (int, error) {
r, _, errno := unix.Syscall6(
unix.SYS_RECVMMSG,
fd,
uintptr(u.sysFd),
uintptr(unsafe.Pointer(&msgs[0])),
uintptr(len(msgs)),
unix.MSG_WAITFORONE,
0,
0,
)
if errno == syscall.EAGAIN || errno == syscall.EWOULDBLOCK {
// No data available, block for I/O and try again.
return int(n), false, nil
}
if errno != 0 {
return int(n), true, &net.OpError{Op: "recvmmsg", Err: errno}
}
return int(n), true, nil
}
func (u *StdConn) listenOutSingle(r EncReader) error {
var err error
var n int
var from netip.AddrPort
buffer := make([]byte, MTU)
for {
n, from, err = u.udpConn.ReadFromUDPAddrPort(buffer)
if err != nil {
return err
if u.closed.Load() {
return 0, net.ErrClosed
}
from = netip.AddrPortFrom(from.Addr().Unmap(), from.Port())
r(from, buffer[:n])
return 0, &net.OpError{Op: "recvmmsg", Err: errno}
}
n := int(r)
if (n == 0 || msgs[0].Len == 0) && u.closed.Load() {
return 0, net.ErrClosed
}
return n, nil
}
func (u *StdConn) listenOutBatch(r EncReader) error {
// recvmsg does one blocking recvmsg into msgs[0]
func (u *StdConn) recvmsg(msgs []rawMessage) (int, error) {
r, _, errno := unix.Syscall6(
unix.SYS_RECVMSG,
uintptr(u.sysFd),
uintptr(unsafe.Pointer(&msgs[0].Hdr)),
0,
0,
0,
0,
)
if errno != 0 {
if u.closed.Load() {
return 0, net.ErrClosed
}
return 0, &net.OpError{Op: "recvmsg", Err: errno}
}
if r == 0 && u.closed.Load() {
return 0, net.ErrClosed
}
msgs[0].Len = uint32(r)
return 1, nil
}
func (u *StdConn) ListenOut(r EncReader) error {
var ip netip.Addr
var n int
var operr error
msgs, buffers, names := u.PrepareRawMessages(u.batch)
//reader needs to capture variables from this function, since it's used as a lambda with rawConn.Read
//defining it outside the loop so it gets re-used
reader := func(fd uintptr) (done bool) {
n, done, operr = recvmmsg(fd, msgs)
return done
read := u.recvmmsg
if u.batch == 1 {
read = u.recvmsg
}
for {
err := u.rawConn.Read(reader)
n, err := read(msgs)
if err != nil {
if errors.Is(err, unix.EINTR) {
continue // interrupted by a signal, retry the read
}
// net.ErrClosed after Close() is teardown, absorbed by the caller's
// closed flag like the other platforms; anything else is a real error.
return err
}
if operr != nil {
return operr
}
for i := 0; i < n; i++ {
// Its ok to skip the ok check here, the slicing is the only error that can occur and it will panic
@@ -222,26 +193,68 @@ func (u *StdConn) listenOutBatch(r EncReader) error {
}
}
func (u *StdConn) ListenOut(r EncReader) error {
if u.batch == 1 {
return u.listenOutSingle(r)
} else {
return u.listenOutBatch(r)
func (u *StdConn) WriteTo(b []byte, ip netip.AddrPort) error {
if u.isV4 {
return u.writeTo4(b, ip)
}
return u.writeTo6(b, ip)
}
func (u *StdConn) writeTo6(b []byte, ip netip.AddrPort) error {
var rsa unix.RawSockaddrInet6
rsa.Family = unix.AF_INET6
rsa.Addr = ip.Addr().As16()
binary.BigEndian.PutUint16((*[2]byte)(unsafe.Pointer(&rsa.Port))[:], ip.Port())
for {
_, _, err := unix.Syscall6(
unix.SYS_SENDTO,
uintptr(u.sysFd),
uintptr(unsafe.Pointer(&b[0])),
uintptr(len(b)),
uintptr(0),
uintptr(unsafe.Pointer(&rsa)),
uintptr(unix.SizeofSockaddrInet6),
)
if err != 0 {
return &net.OpError{Op: "sendto", Err: err}
}
return nil
}
}
func (u *StdConn) WriteTo(b []byte, ip netip.AddrPort) error {
_, err := u.udpConn.WriteToUDPAddrPort(b, ip)
return err
func (u *StdConn) writeTo4(b []byte, ip netip.AddrPort) error {
if !ip.Addr().Is4() {
return ErrInvalidIPv6RemoteForSocket
}
var rsa unix.RawSockaddrInet4
rsa.Family = unix.AF_INET
rsa.Addr = ip.Addr().As4()
binary.BigEndian.PutUint16((*[2]byte)(unsafe.Pointer(&rsa.Port))[:], ip.Port())
for {
_, _, err := unix.Syscall6(
unix.SYS_SENDTO,
uintptr(u.sysFd),
uintptr(unsafe.Pointer(&b[0])),
uintptr(len(b)),
uintptr(0),
uintptr(unsafe.Pointer(&rsa)),
uintptr(unix.SizeofSockaddrInet4),
)
if err != 0 {
return &net.OpError{Op: "sendto", Err: err}
}
return nil
}
}
func (u *StdConn) ReloadConfig(c *config.C) {
b := c.GetInt("listen.read_buffer", 0)
if b > 0 {
err := u.SetRecvBuffer(b)
if err == nil {
s, err := u.GetRecvBuffer()
if err == nil {
if err := u.SetRecvBuffer(b); err == nil {
if s, err := u.GetRecvBuffer(); err == nil {
u.l.Info("listen.read_buffer was set", "size", s)
} else {
u.l.Warn("Failed to get listen.read_buffer", "error", err)
@@ -253,10 +266,8 @@ func (u *StdConn) ReloadConfig(c *config.C) {
b = c.GetInt("listen.write_buffer", 0)
if b > 0 {
err := u.SetSendBuffer(b)
if err == nil {
s, err := u.GetSendBuffer()
if err == nil {
if err := u.SetSendBuffer(b); err == nil {
if s, err := u.GetSendBuffer(); err == nil {
u.l.Info("listen.write_buffer was set", "size", s)
} else {
u.l.Warn("Failed to get listen.write_buffer", "error", err)
@@ -269,10 +280,8 @@ func (u *StdConn) ReloadConfig(c *config.C) {
b = c.GetInt("listen.so_mark", 0)
s, err := u.GetSoMark()
if b > 0 || (err == nil && s != 0) {
err := u.SetSoMark(b)
if err == nil {
s, err := u.GetSoMark()
if err == nil {
if err := u.SetSoMark(b); err == nil {
if s, err := u.GetSoMark(); err == nil {
u.l.Info("listen.so_mark was set", "mark", s)
} else {
u.l.Warn("Failed to get listen.so_mark", "error", err)
@@ -285,28 +294,20 @@ func (u *StdConn) ReloadConfig(c *config.C) {
func (u *StdConn) getMemInfo(meminfo *[unix.SK_MEMINFO_VARS]uint32) error {
var vallen uint32 = 4 * unix.SK_MEMINFO_VARS
if u.rawConn == nil {
return fmt.Errorf("no UDP connection")
}
var opErr error
err := u.rawConn.Control(func(fd uintptr) {
_, _, syserr := unix.Syscall6(unix.SYS_GETSOCKOPT, fd, uintptr(unix.SOL_SOCKET), uintptr(unix.SO_MEMINFO), uintptr(unsafe.Pointer(meminfo)), uintptr(unsafe.Pointer(&vallen)), 0)
if syserr != 0 {
opErr = syserr
}
})
if err != nil {
_, _, err := unix.Syscall6(unix.SYS_GETSOCKOPT, uintptr(u.sysFd), uintptr(unix.SOL_SOCKET), uintptr(unix.SO_MEMINFO), uintptr(unsafe.Pointer(meminfo)), uintptr(unsafe.Pointer(&vallen)), 0)
if err != 0 {
return err
}
return opErr
return nil
}
func (u *StdConn) Close() error {
if u.udpConn != nil {
return u.udpConn.Close()
}
return nil
u.closed.Store(true)
// Wake the reader parked in recvmmsg/recvmsg. shutdown(2) on an unconnected socket
// returns ENOTCONN but still wakes it, so ignore the error.
// The reader then sees closed and stops touching the fd, making the Close below safe.
_ = unix.Shutdown(u.sysFd, unix.SHUT_RDWR)
return unix.Close(u.sysFd)
}
func NewUDPStatsEmitter(udpConns []Conn) func() {
+179
View File
@@ -0,0 +1,179 @@
//go:build linux && !android && !e2e_testing
package udp
import (
"errors"
"fmt"
"log/slog"
"net"
"net/netip"
"os"
"runtime"
"sync/atomic"
"testing"
"time"
"golang.org/x/sys/unix"
)
func testLogger() *slog.Logger {
return slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: slog.LevelError}))
}
// TestShutdownWakesAfterRx_Mechanism exercises the kernel quirk our teardown
// relies on: once a socket has received a packet, shutdown(2) wakes a blocked
// recvmmsg with n>=1/Len==0 (not n==0). recvmmsg must turn that into net.ErrClosed
// once Close set closed, so a parked reader exits instead of spinning.
func TestShutdownWakesAfterRx_Mechanism(t *testing.T) {
c, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, true, 64)
if err != nil {
t.Fatalf("NewListener: %v", err)
}
sc := c.(*StdConn)
addr, err := sc.LocalAddr()
if err != nil {
t.Fatalf("LocalAddr: %v", err)
}
msgs, _, _ := sc.PrepareRawMessages(sc.batch)
// Receive a real packet so the socket has carried data.
send, err := net.Dial("udp", addr.String())
if err != nil {
t.Fatalf("dial: %v", err)
}
if _, err := send.Write([]byte("hello")); err != nil {
t.Fatalf("write: %v", err)
}
time.Sleep(50 * time.Millisecond)
n, err := sc.recvmmsg(msgs)
t.Logf("drain of real packet: n=%d err=%v msgs[0].Len=%d", n, err, msgs[0].Len)
_ = send.Close()
// Block a reader on the now-empty queue, then tear down as Close() does.
// recvmmsg must return net.ErrClosed (not hang, not spin) even post-rx.
done := make(chan error, 1)
go func() {
_, err := sc.recvmmsg(msgs)
done <- err
}()
time.Sleep(150 * time.Millisecond) // let it park in recvmmsg
sc.closed.Store(true)
if serr := unix.Shutdown(sc.sysFd, unix.SHUT_RDWR); serr != nil {
t.Logf("shutdown returned %v (expected ENOTCONN on unconnected UDP)", serr)
}
select {
case err := <-done:
if !errors.Is(err, net.ErrClosed) {
t.Errorf("recvmmsg after post-rx shutdown returned %v, want net.ErrClosed", err)
}
case <-time.After(2 * time.Second):
t.Fatalf("HANG: recvmmsg did not return after shutdown following a received packet")
}
_ = unix.Close(sc.sysFd)
}
// TestListenOutTeardown_TrafficPatterns reproduces the field report: a blocking
// reader must tear down cleanly on Close() regardless of what the socket has
// carried. The three cases the report called out:
//
// no traffic ever -> works (shutdown wakes recvmmsg with n==0)
// ping once, then idle -> historically HUNG: once the socket has received a
// packet, shutdown(2) wakes recvmmsg with n>=1/Len==0,
// which an n==0-only teardown check misses
// continuous traffic -> works (a real packet is always arriving)
//
// All three must return within the deadline; a hang dumps goroutines so the
// stuck reader is visible.
func TestListenOutTeardown_TrafficPatterns(t *testing.T) {
cases := []struct {
name string
traffic func(send net.Conn, stop <-chan struct{})
}{
{"no_traffic_ever", func(net.Conn, <-chan struct{}) {}},
{"ping_once_then_idle", func(send net.Conn, _ <-chan struct{}) {
_, _ = send.Write([]byte("hello"))
}},
{"continuous", func(send net.Conn, stop <-chan struct{}) {
for {
select {
case <-stop:
return
default:
_, _ = send.Write([]byte("hello"))
time.Sleep(2 * time.Millisecond)
}
}
}},
}
// batch 1 exercises the recvmsg path, batch 64 the recvmmsg path; both must
// tear down cleanly.
for _, batch := range []int{1, 64} {
for _, tc := range cases {
t.Run(fmt.Sprintf("batch%d/%s", batch, tc.name), func(t *testing.T) {
runTeardownCase(t, batch, tc.name, tc.traffic)
})
}
}
}
func runTeardownCase(t *testing.T, batch int, name string, traffic func(send net.Conn, stop <-chan struct{})) {
c, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, true, batch)
if err != nil {
t.Fatalf("NewListener: %v", err)
}
sc := c.(*StdConn)
addr, err := sc.LocalAddr()
if err != nil {
t.Fatalf("LocalAddr: %v", err)
}
var received atomic.Int64
loopDone := make(chan error, 1)
go func() {
loopDone <- sc.ListenOut(func(netip.AddrPort, []byte) {
received.Add(1)
})
}()
send, err := net.Dial("udp", addr.String())
if err != nil {
t.Fatalf("dial: %v", err)
}
defer send.Close()
stop := make(chan struct{})
trafficDone := make(chan struct{})
go func() {
traffic(send, stop)
close(trafficDone)
}()
// Let the pattern run and, for the idle case, the reader park again on an
// empty queue with the socket already having received a packet.
time.Sleep(500 * time.Millisecond)
start := time.Now()
if err := sc.Close(); err != nil {
t.Fatalf("Close: %v", err)
}
close(stop)
select {
case err := <-loopDone:
// Clean teardown surfaces as net.ErrClosed (propagated like the other
// platforms); the caller absorbs it via its closed flag.
if err != nil && !errors.Is(err, net.ErrClosed) {
t.Fatalf("%s: ListenOut returned unexpected error on teardown: %v", name, err)
}
t.Logf("%s: closed in %v (received %d packets)", name, time.Since(start), received.Load())
case <-time.After(3 * time.Second):
buf := make([]byte, 1<<20)
n := runtime.Stack(buf, true)
t.Fatalf("%s: HANG, ListenOut did not return within 3s of Close\n%s", name, buf[:n])
}
<-trafficDone
}
+20 -6
View File
@@ -10,6 +10,7 @@ import (
"net/netip"
"os"
"sync"
"sync/atomic"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/header"
@@ -64,7 +65,9 @@ func acquirePacket() *Packet {
}
type TesterConn struct {
Addr netip.AddrPort
// addr is read by nebula's own goroutines on every send and by the router's flow renderer, and a test can
// move it mid-run to simulate roaming, so it is atomic rather than a plain field.
addr atomic.Pointer[netip.AddrPort]
RxPackets chan *Packet // Packets to receive into nebula
TxPackets chan *Packet // Packets transmitted outside by nebula
@@ -82,13 +85,24 @@ type TesterConn struct {
}
func NewListener(l *slog.Logger, ip netip.Addr, port int, _ bool, _ int) (Conn, error) {
return &TesterConn{
Addr: netip.AddrPortFrom(ip, uint16(port)),
c := &TesterConn{
RxPackets: make(chan *Packet, 10),
TxPackets: make(chan *Packet, 10),
done: make(chan struct{}),
l: l,
}, nil
}
c.SetAddr(netip.AddrPortFrom(ip, uint16(port)))
return c, nil
}
// GetAddr returns the underlay address this conn currently sends from.
func (u *TesterConn) GetAddr() netip.AddrPort {
return *u.addr.Load()
}
// SetAddr moves this conn to a new underlay address, standing in for a host waking up on a different network.
func (u *TesterConn) SetAddr(addr netip.AddrPort) {
u.addr.Store(&addr)
}
// Send will place a UdpPacket onto the receive queue for nebula to consume
@@ -147,7 +161,7 @@ func (u *TesterConn) WriteTo(b []byte, addr netip.AddrPort) error {
p.Data = p.Data[:len(b)]
}
copy(p.Data, b)
p.From = u.Addr
p.From = u.GetAddr()
p.To = addr
select {
case <-u.done:
@@ -178,7 +192,7 @@ func NewUDPStatsEmitter(_ []Conn) func() {
}
func (u *TesterConn) LocalAddr() (netip.AddrPort, error) {
return u.Addr, nil
return u.GetAddr(), nil
}
func (u *TesterConn) SupportsMultipleReaders() bool {