Compare commits

..

4 Commits

Author SHA1 Message Date
JackDoan 2c30f016a4 yey 2026-04-15 16:15:19 -05:00
JackDoan 4c2ea20e0c goddamnit 2026-04-15 15:06:02 -05:00
JackDoan b7466f0d7d goddamnit 2026-04-15 14:58:53 -05:00
JackDoan ffa1f4994f tun_linux.go: set O_NONBLOCK 2026-04-15 14:30:16 -05:00
13 changed files with 94 additions and 167 deletions
+2 -10
View File
@@ -78,16 +78,8 @@ func main() {
} }
if !*configTest { if !*configTest {
wait, err := ctrl.Start() ctrl.Start()
if err != nil { ctrl.ShutdownBlock()
util.LogWithContextIfNeeded("Error while running", err, l)
os.Exit(1)
}
go ctrl.ShutdownBlock()
wait()
l.Info("Goodbye")
} }
os.Exit(0) os.Exit(0)
+2 -10
View File
@@ -72,17 +72,9 @@ func main() {
} }
if !*configTest { if !*configTest {
wait, err := ctrl.Start() ctrl.Start()
if err != nil {
util.LogWithContextIfNeeded("Error while running", err, l)
os.Exit(1)
}
go ctrl.ShutdownBlock()
notifyReady(l) notifyReady(l)
wait() ctrl.ShutdownBlock()
l.Info("Goodbye")
} }
os.Exit(0) os.Exit(0)
+6 -52
View File
@@ -2,11 +2,9 @@ package nebula
import ( import (
"context" "context"
"errors"
"net/netip" "net/netip"
"os" "os"
"os/signal" "os/signal"
"sync"
"syscall" "syscall"
"github.com/sirupsen/logrus" "github.com/sirupsen/logrus"
@@ -15,16 +13,6 @@ import (
"github.com/slackhq/nebula/overlay" "github.com/slackhq/nebula/overlay"
) )
type RunState int
const (
Stopped RunState = 0 // The control has yet to be started
Started RunState = 1 // The control has been started
Stopping RunState = 2 // The control is stopping
)
var ErrAlreadyStarted = errors.New("nebula is already started")
// Every interaction here needs to take extra care to copy memory and not return or use arguments "as is" when touching // Every interaction here needs to take extra care to copy memory and not return or use arguments "as is" when touching
// core. This means copying IP objects, slices, de-referencing pointers and taking the actual value, etc // core. This means copying IP objects, slices, de-referencing pointers and taking the actual value, etc
@@ -38,9 +26,6 @@ type controlHostLister interface {
} }
type Control struct { type Control struct {
stateLock sync.Mutex
state RunState
f *Interface f *Interface
l *logrus.Logger l *logrus.Logger
ctx context.Context ctx context.Context
@@ -64,21 +49,10 @@ type ControlHostInfo struct {
CurrentRelaysThroughMe []netip.Addr `json:"currentRelaysThroughMe"` CurrentRelaysThroughMe []netip.Addr `json:"currentRelaysThroughMe"`
} }
// Start actually runs nebula, this is a nonblocking call. // Start actually runs nebula, this is a nonblocking call. To block use Control.ShutdownBlock()
// The returned function can be used to wait for nebula to fully stop. func (c *Control) Start() {
func (c *Control) Start() (func(), error) {
c.stateLock.Lock()
if c.state != Stopped {
c.stateLock.Unlock()
return nil, ErrAlreadyStarted
}
// Activate the interface // Activate the interface
err := c.f.activate() c.f.activate()
if err != nil {
c.stateLock.Unlock()
return nil, err
}
// Call all the delayed funcs that waited patiently for the interface to be created. // Call all the delayed funcs that waited patiently for the interface to be created.
if c.sshStart != nil { if c.sshStart != nil {
@@ -98,33 +72,15 @@ func (c *Control) Start() (func(), error) {
} }
// Start reading packets. // Start reading packets.
c.state = Started c.f.run()
c.stateLock.Unlock()
return c.f.run()
}
func (c *Control) State() RunState {
c.stateLock.Lock()
defer c.stateLock.Unlock()
return c.state
} }
func (c *Control) Context() context.Context { func (c *Control) Context() context.Context {
return c.ctx return c.ctx
} }
// Stop is a non-blocking call that signals nebula to close all tunnels and shut down // Stop signals nebula to shutdown and close all tunnels, returns after the shutdown is complete
func (c *Control) Stop() { func (c *Control) Stop() {
c.stateLock.Lock()
if c.state != Started {
c.stateLock.Unlock()
// We are stopping or stopped already
return
}
c.state = Stopping
c.stateLock.Unlock()
// Stop the handshakeManager (and other services), to prevent new tunnels from // Stop the handshakeManager (and other services), to prevent new tunnels from
// being created while we're shutting them all down. // being created while we're shutting them all down.
c.cancel() c.cancel()
@@ -133,9 +89,7 @@ func (c *Control) Stop() {
if err := c.f.Close(); err != nil { if err := c.f.Close(); err != nil {
c.l.WithError(err).Error("Close interface failed") c.l.WithError(err).Error("Close interface failed")
} }
c.stateLock.Lock() c.l.Info("Goodbye")
c.state = Stopped
c.stateLock.Unlock()
} }
// ShutdownBlock will listen for and block on term and interrupt signals, calling Control.Stop() once signalled // ShutdownBlock will listen for and block on term and interrupt signals, calling Control.Stop() once signalled
+22 -41
View File
@@ -6,7 +6,7 @@ import (
"fmt" "fmt"
"io" "io"
"net/netip" "net/netip"
"sync" "os"
"sync/atomic" "sync/atomic"
"time" "time"
@@ -87,7 +87,6 @@ type Interface struct {
writers []udp.Conn writers []udp.Conn
readers []io.ReadWriteCloser readers []io.ReadWriteCloser
wg sync.WaitGroup
metricHandshakes metrics.Histogram metricHandshakes metrics.Histogram
messageMetrics *MessageMetrics messageMetrics *MessageMetrics
@@ -210,7 +209,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
// activate creates the interface on the host. After the interface is created, any // activate creates the interface on the host. After the interface is created, any
// other services that want to bind listeners to its IP may do so successfully. However, // other services that want to bind listeners to its IP may do so successfully. However,
// the interface isn't going to process anything until run() is called. // the interface isn't going to process anything until run() is called.
func (f *Interface) activate() error { func (f *Interface) activate() {
// actually turn on tun dev // actually turn on tun dev
addr, err := f.outside.LocalAddr() addr, err := f.outside.LocalAddr()
@@ -238,36 +237,28 @@ func (f *Interface) activate() error {
if i > 0 { if i > 0 {
reader, err = f.inside.NewMultiQueueReader() reader, err = f.inside.NewMultiQueueReader()
if err != nil { if err != nil {
return err f.l.Fatal(err)
} }
} }
f.readers[i] = reader f.readers[i] = reader
} }
if err = f.inside.Activate(); err != nil { if err := f.inside.Activate(); err != nil {
f.inside.Close() f.inside.Close()
return err f.l.Fatal(err)
} }
return nil
} }
func (f *Interface) run() (func(), error) { func (f *Interface) run() {
// Launch n queues to read packets from udp // Launch n queues to read packets from udp
for i := 0; i < f.routines; i++ { for i := 0; i < f.routines; i++ {
f.wg.Go(func() { go f.listenOut(i)
f.listenOut(i)
})
} }
// Launch n queues to read packets from tun dev // Launch n queues to read packets from tun dev
for i := 0; i < f.routines; i++ { for i := 0; i < f.routines; i++ {
f.wg.Go(func() { go f.listenIn(f.readers[i], i)
f.listenIn(f.readers[i], i)
})
} }
return f.wg.Wait, nil
} }
func (f *Interface) listenOut(i int) { func (f *Interface) listenOut(i int) {
@@ -285,16 +276,9 @@ func (f *Interface) listenOut(i int) {
fwPacket := &firewall.Packet{} fwPacket := &firewall.Packet{}
nb := make([]byte, 12, 12) nb := make([]byte, 12, 12)
err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) { li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get(f.l)) f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get(f.l))
}) })
if err != nil && !f.closed.Load() {
f.l.WithError(err).Error("Error while reading inbound packet, closing")
//TODO: Trigger Control to close
}
f.l.Infof("underlay reader %v is done", i)
} }
func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) { func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
@@ -308,17 +292,17 @@ func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
for { for {
n, err := reader.Read(packet) n, err := reader.Read(packet)
if err != nil { if err != nil {
if !f.closed.Load() { if errors.Is(err, os.ErrClosed) && f.closed.Load() {
f.l.WithError(err).WithField("reader", i).Error("Error while reading outbound packet, closing") return
//TODO: Trigger Control to close
} }
break
f.l.WithError(err).Error("Error while reading outbound packet")
// This only seems to happen when something fatal happens to the fd, so exit.
os.Exit(2)
} }
f.consumeInsidePacket(packet[:n], fwPacket, nb, out, i, conntrackCache.Get(f.l)) f.consumeInsidePacket(packet[:n], fwPacket, nb, out, i, conntrackCache.Get(f.l))
} }
f.l.Infof("overlay reader %v is done", i)
} }
func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) { func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
@@ -493,26 +477,23 @@ func (f *Interface) GetCertState() *CertState {
} }
func (f *Interface) Close() error { func (f *Interface) Close() error {
var err error
f.closed.Store(true) f.closed.Store(true)
// Release the udp readers for _, u := range f.writers {
for i, u := range f.writers { err := u.Close()
err = u.Close()
if err != nil { if err != nil {
f.l.WithError(err).WithField("writer", i).Error("Error while closing udp socket") f.l.WithError(err).Error("Error while closing udp socket")
} }
} }
// Release the tun readers
for i, r := range f.readers { for i, r := range f.readers {
if i == 0 { if i == 0 {
continue // f.readers[0] is f.inside, which we want to save for last, since it closes other stuff too continue // f.readers[0] is f.inside, which we want to save for last
} }
if err = r.Close(); err != nil { if err := r.Close(); err != nil {
f.l.WithError(err).WithField("reader", i).Error("Error while closing tun reader") f.l.WithError(err).Error("Error while closing tun reader")
} }
} }
// Release the tun device
return f.inside.Close() return f.inside.Close()
} }
+9 -9
View File
@@ -288,15 +288,15 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
} }
return &Control{ return &Control{
f: ifce, ifce,
l: l, l,
ctx: ctx, ctx,
cancel: cancel, cancel,
sshStart: sshStart, sshStart,
statsStart: statsStart, statsStart,
dnsStart: dnsStart, dnsStart,
lighthouseStart: lightHouse.StartUpdateWorker, lightHouse.StartUpdateWorker,
connectionManagerStart: connManager.Start, connManager.Start,
}, nil }, nil
} }
+12 -15
View File
@@ -72,6 +72,11 @@ type ifreqQLEN struct {
func newTunFromFd(c *config.C, l *logrus.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) { func newTunFromFd(c *config.C, l *logrus.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
file := os.NewFile(uintptr(deviceFd), "/dev/net/tun") file := os.NewFile(uintptr(deviceFd), "/dev/net/tun")
err := unix.SetNonblock(deviceFd, true)
if err != nil {
_ = file.Close()
return nil, err
}
t, err := newTunGeneric(c, l, file, vpnNetworks) t, err := newTunGeneric(c, l, file, vpnNetworks)
if err != nil { if err != nil {
@@ -237,7 +242,7 @@ func (t *tun) SupportsMultiqueue() bool {
} }
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) { func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0) fd, err := unix.Open("/dev/net/tun", os.O_RDWR|unix.O_NONBLOCK, 0)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -684,23 +689,15 @@ func (t *tun) updateRoutes(r netlink.RouteUpdate) {
func (t *tun) Close() error { func (t *tun) Close() error {
if t.routeChan != nil { if t.routeChan != nil {
close(t.routeChan) close(t.routeChan)
t.routeChan = nil
}
if t.ioctlFd > 0 {
err := os.NewFile(t.ioctlFd, "ioctlFd").Close()
if err != nil {
t.l.WithField("error", err).Error("Failed to close ioctl fd")
}
t.ioctlFd = 0
} }
if t.ReadWriteCloser != nil { if t.ReadWriteCloser != nil {
err := t.ReadWriteCloser.Close() _ = t.ReadWriteCloser.Close()
if err != nil { }
t.l.WithField("error", err).Error("Failed to close tun file")
return err if t.ioctlFd > 0 {
} _ = os.NewFile(t.ioctlFd, "ioctlFd").Close()
t.ioctlFd = 0
} }
return nil return nil
+1 -10
View File
@@ -44,10 +44,7 @@ type Service struct {
} }
func New(control *nebula.Control) (*Service, error) { func New(control *nebula.Control) (*Service, error) {
wait, err := control.Start() control.Start()
if err != nil {
return nil, err
}
ctx := control.Context() ctx := control.Context()
eg, ctx := errgroup.WithContext(ctx) eg, ctx := errgroup.WithContext(ctx)
@@ -144,12 +141,6 @@ func New(control *nebula.Control) (*Service, error) {
} }
}) })
// Add the nebula wait function to the group
eg.Go(func() error {
wait()
return nil
})
return &s, nil return &s, nil
} }
+3 -3
View File
@@ -16,7 +16,7 @@ type EncReader func(
type Conn interface { type Conn interface {
Rebind() error Rebind() error
LocalAddr() (netip.AddrPort, error) LocalAddr() (netip.AddrPort, error)
ListenOut(r EncReader) error ListenOut(r EncReader)
WriteTo(b []byte, addr netip.AddrPort) error WriteTo(b []byte, addr netip.AddrPort) error
ReloadConfig(c *config.C) ReloadConfig(c *config.C)
SupportsMultipleReaders() bool SupportsMultipleReaders() bool
@@ -31,8 +31,8 @@ func (NoopConn) Rebind() error {
func (NoopConn) LocalAddr() (netip.AddrPort, error) { func (NoopConn) LocalAddr() (netip.AddrPort, error) {
return netip.AddrPort{}, nil return netip.AddrPort{}, nil
} }
func (NoopConn) ListenOut(_ EncReader) error { func (NoopConn) ListenOut(_ EncReader) {
return nil return
} }
func (NoopConn) SupportsMultipleReaders() bool { func (NoopConn) SupportsMultipleReaders() bool {
return false return false
+3 -2
View File
@@ -165,7 +165,7 @@ func NewUDPStatsEmitter(udpConns []Conn) func() {
return func() {} return func() {}
} }
func (u *StdConn) ListenOut(r EncReader) error { func (u *StdConn) ListenOut(r EncReader) {
buffer := make([]byte, MTU) buffer := make([]byte, MTU)
for { for {
@@ -173,7 +173,8 @@ func (u *StdConn) ListenOut(r EncReader) error {
n, rua, err := u.ReadFromUDPAddrPort(buffer) n, rua, err := u.ReadFromUDPAddrPort(buffer)
if err != nil { if err != nil {
if errors.Is(err, net.ErrClosed) { if errors.Is(err, net.ErrClosed) {
return err u.l.WithError(err).Debug("udp socket is closed, exiting read loop")
return
} }
u.l.WithError(err).Error("unexpected udp socket receive error") u.l.WithError(err).Error("unexpected udp socket receive error")
+15 -2
View File
@@ -10,9 +10,11 @@ package udp
import ( import (
"context" "context"
"errors"
"fmt" "fmt"
"net" "net"
"net/netip" "net/netip"
"time"
"github.com/sirupsen/logrus" "github.com/sirupsen/logrus"
"github.com/slackhq/nebula/config" "github.com/slackhq/nebula/config"
@@ -71,14 +73,25 @@ type rawMessage struct {
Len uint32 Len uint32
} }
func (u *GenericConn) ListenOut(r EncReader) error { func (u *GenericConn) ListenOut(r EncReader) {
buffer := make([]byte, MTU) buffer := make([]byte, MTU)
var lastRecvErr time.Time
for { for {
// Just read one packet at a time // Just read one packet at a time
n, rua, err := u.ReadFromUDPAddrPort(buffer) n, rua, err := u.ReadFromUDPAddrPort(buffer)
if err != nil { if err != nil {
return err if errors.Is(err, net.ErrClosed) {
u.l.WithError(err).Debug("udp socket is closed, exiting read loop")
return
}
// Dampen unexpected message warns to once per minute
if lastRecvErr.IsZero() || time.Since(lastRecvErr) > time.Minute {
lastRecvErr = time.Now()
u.l.WithError(err).Warn("unexpected udp socket receive error")
}
continue
} }
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n]) r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n])
+14 -8
View File
@@ -171,7 +171,7 @@ func recvmmsg(fd uintptr, msgs []rawMessage) (int, bool, error) {
return int(n), true, nil return int(n), true, nil
} }
func (u *StdConn) listenOutSingle(r EncReader) error { func (u *StdConn) listenOutSingle(r EncReader) {
var err error var err error
var n int var n int
var from netip.AddrPort var from netip.AddrPort
@@ -180,14 +180,15 @@ func (u *StdConn) listenOutSingle(r EncReader) error {
for { for {
n, from, err = u.udpConn.ReadFromUDPAddrPort(buffer) n, from, err = u.udpConn.ReadFromUDPAddrPort(buffer)
if err != nil { if err != nil {
return err u.l.WithError(err).Debug("udp socket is closed, exiting read loop")
return
} }
from = netip.AddrPortFrom(from.Addr().Unmap(), from.Port()) from = netip.AddrPortFrom(from.Addr().Unmap(), from.Port())
r(from, buffer[:n]) r(from, buffer[:n])
} }
} }
func (u *StdConn) listenOutBatch(r EncReader) error { func (u *StdConn) listenOutBatch(r EncReader) {
var ip netip.Addr var ip netip.Addr
var n int var n int
var operr error var operr error
@@ -204,10 +205,12 @@ func (u *StdConn) listenOutBatch(r EncReader) error {
for { for {
err := u.rawConn.Read(reader) err := u.rawConn.Read(reader)
if err != nil { if err != nil {
return err u.l.WithError(err).Debug("udp socket is closed, exiting read loop")
return
} }
if operr != nil { if operr != nil {
return operr u.l.WithError(operr).Debug("operr: udp socket is closed, exiting read loop")
return
} }
for i := 0; i < n; i++ { for i := 0; i < n; i++ {
@@ -222,11 +225,14 @@ func (u *StdConn) listenOutBatch(r EncReader) error {
} }
} }
func (u *StdConn) ListenOut(r EncReader) error { func (u *StdConn) ListenOut(r EncReader) {
if u.batch == 1 { if u.batch == 1 {
return u.listenOutSingle(r) //save some ram by not calling PrepareRawMessages for fields we won't use
//we could also make this path more common by calling recvmmsg with msgs[:1],
//but that's still the recvmmsg syscall, which would be a change
u.listenOutSingle(r)
} else { } else {
return u.listenOutBatch(r) u.listenOutBatch(r)
} }
} }
+3 -2
View File
@@ -140,7 +140,7 @@ func (u *RIOConn) bind(l *logrus.Logger, sa windows.Sockaddr) error {
return nil return nil
} }
func (u *RIOConn) ListenOut(r EncReader) error { func (u *RIOConn) ListenOut(r EncReader) {
buffer := make([]byte, MTU) buffer := make([]byte, MTU)
var lastRecvErr time.Time var lastRecvErr time.Time
@@ -151,7 +151,8 @@ func (u *RIOConn) ListenOut(r EncReader) error {
if err != nil { if err != nil {
if errors.Is(err, net.ErrClosed) { if errors.Is(err, net.ErrClosed) {
return err u.l.WithError(err).Debug("udp socket is closed, exiting read loop")
return
} }
// Dampen unexpected message warns to once per minute // Dampen unexpected message warns to once per minute
if lastRecvErr.IsZero() || time.Since(lastRecvErr) > time.Minute { if lastRecvErr.IsZero() || time.Since(lastRecvErr) > time.Minute {
+2 -3
View File
@@ -6,7 +6,6 @@ package udp
import ( import (
"io" "io"
"net/netip" "net/netip"
"os"
"sync/atomic" "sync/atomic"
"github.com/sirupsen/logrus" "github.com/sirupsen/logrus"
@@ -107,11 +106,11 @@ func (u *TesterConn) WriteTo(b []byte, addr netip.AddrPort) error {
return nil return nil
} }
func (u *TesterConn) ListenOut(r EncReader) error { func (u *TesterConn) ListenOut(r EncReader) {
for { for {
p, ok := <-u.RxPackets p, ok := <-u.RxPackets
if !ok { if !ok {
return os.ErrClosed return
} }
r(p.From, p.Data) r(p.From, p.Data)
} }