mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 18:17:02 +02:00
Compare commits
6 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 459cfc6f83 | |||
| 86733864fe | |||
| ab736e4c6b | |||
| 5ecdd4eaa9 | |||
| 1b84bd0050 | |||
| 384610f81a |
@@ -53,7 +53,12 @@ func main() {
|
|||||||
l := logging.NewLogger(os.Stdout)
|
l := logging.NewLogger(os.Stdout)
|
||||||
|
|
||||||
if *serviceFlag != "" {
|
if *serviceFlag != "" {
|
||||||
if err := doService(configPath, configTest, Build, serviceFlag); err != nil {
|
if *configTest {
|
||||||
|
fmt.Println("-test is not supported with -service, run the config test without -service")
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := doService(configPath, Build, serviceFlag); err != nil {
|
||||||
l.Error("Service command failed", "error", err)
|
l.Error("Service command failed", "error", err)
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
@@ -93,15 +98,14 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !*configTest {
|
if !*configTest {
|
||||||
wait, err := ctrl.Start()
|
if err := ctrl.Start(); err != nil {
|
||||||
if err != nil {
|
|
||||||
util.LogWithContextIfNeeded("Error while running", err, l)
|
util.LogWithContextIfNeeded("Error while running", err, l)
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
go ctrl.ShutdownBlock()
|
go ctrl.ShutdownBlock()
|
||||||
|
|
||||||
if err := wait(); err != nil {
|
if err := ctrl.Wait(); err != nil {
|
||||||
l.Error("Nebula stopped due to fatal error", "error", err)
|
l.Error("Nebula stopped due to fatal error", "error", err)
|
||||||
os.Exit(2)
|
os.Exit(2)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package main
|
|||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
"log"
|
||||||
|
"os"
|
||||||
|
|
||||||
"github.com/kardianos/service"
|
"github.com/kardianos/service"
|
||||||
"github.com/slackhq/nebula"
|
"github.com/slackhq/nebula"
|
||||||
@@ -14,7 +15,6 @@ var logger service.Logger
|
|||||||
|
|
||||||
type program struct {
|
type program struct {
|
||||||
configPath *string
|
configPath *string
|
||||||
configTest *bool
|
|
||||||
build string
|
build string
|
||||||
control *nebula.Control
|
control *nebula.Control
|
||||||
}
|
}
|
||||||
@@ -40,22 +40,41 @@ func (p *program) Start(s service.Service) error {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
p.control, err = nebula.Main(c, *p.configTest, Build, l, nil)
|
p.control, err = nebula.Main(c, false, Build, l, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
p.control.Start()
|
if err := p.control.Start(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Nebula can stop itself on a fatal packet reader error, make sure to log it if it happens.
|
||||||
|
go func() {
|
||||||
|
if err := p.control.Wait(); err != nil {
|
||||||
|
logger.Error(fmt.Sprintf("Nebula stopped due to fatal error: %v", err))
|
||||||
|
os.Exit(2)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *program) Stop(s service.Service) error {
|
func (p *program) Stop(s service.Service) error {
|
||||||
logger.Info("Nebula service stopping.")
|
logger.Info("Nebula service stopping.")
|
||||||
|
if p.control == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
p.control.Stop()
|
p.control.Stop()
|
||||||
|
|
||||||
|
// block until nebula has fully drained before reporting stopped.
|
||||||
|
// error logging is handled by Start.
|
||||||
|
_ = p.control.Wait()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func doService(configPath *string, configTest *bool, build string, serviceFlag *string) error {
|
func doService(configPath *string, build string, serviceFlag *string) error {
|
||||||
if *configPath == "" {
|
if *configPath == "" {
|
||||||
p, err := config.DefaultPath()
|
p, err := config.DefaultPath()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -73,7 +92,6 @@ func doService(configPath *string, configTest *bool, build string, serviceFlag *
|
|||||||
|
|
||||||
prg := &program{
|
prg := &program{
|
||||||
configPath: configPath,
|
configPath: configPath,
|
||||||
configTest: configTest,
|
|
||||||
build: build,
|
build: build,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -105,8 +123,9 @@ func doService(configPath *string, configTest *bool, build string, serviceFlag *
|
|||||||
switch *serviceFlag {
|
switch *serviceFlag {
|
||||||
case "run":
|
case "run":
|
||||||
if err := s.Run(); err != nil {
|
if err := s.Run(); err != nil {
|
||||||
// Route any errors to the system logger
|
// Route any errors to the system logger and report the failure
|
||||||
logger.Error(err)
|
logger.Error(err)
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
default:
|
default:
|
||||||
if err := service.Control(s, *serviceFlag); err != nil {
|
if err := service.Control(s, *serviceFlag); err != nil {
|
||||||
|
|||||||
+2
-3
@@ -84,8 +84,7 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !*configTest {
|
if !*configTest {
|
||||||
wait, err := ctrl.Start()
|
if err := ctrl.Start(); err != nil {
|
||||||
if err != nil {
|
|
||||||
util.LogWithContextIfNeeded("Error while running", err, l)
|
util.LogWithContextIfNeeded("Error while running", err, l)
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
@@ -93,7 +92,7 @@ func main() {
|
|||||||
go ctrl.ShutdownBlock()
|
go ctrl.ShutdownBlock()
|
||||||
notifyReady(l)
|
notifyReady(l)
|
||||||
|
|
||||||
if err := wait(); err != nil {
|
if err := ctrl.Wait(); err != nil {
|
||||||
l.Error("Nebula stopped due to fatal error", "error", err)
|
l.Error("Nebula stopped due to fatal error", "error", err)
|
||||||
os.Exit(2)
|
os.Exit(2)
|
||||||
}
|
}
|
||||||
|
|||||||
+49
-23
@@ -69,29 +69,29 @@ type ControlHostInfo struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Start actually runs nebula, this is a nonblocking call.
|
// Start actually runs nebula, this is a nonblocking call.
|
||||||
// The returned function blocks until nebula has fully stopped and returns the
|
// Use Wait to block until nebula has fully stopped and to learn whether a fatal reader error caused the shutdown.
|
||||||
// first fatal reader error (if any). A nil error means nebula shut down
|
func (c *Control) Start() error {
|
||||||
// gracefully; a non-nil error means a reader hit an unexpected failure that
|
|
||||||
// triggered the shutdown.
|
|
||||||
func (c *Control) Start() (func() error, error) {
|
|
||||||
c.stateLock.Lock()
|
c.stateLock.Lock()
|
||||||
defer c.stateLock.Unlock()
|
defer c.stateLock.Unlock()
|
||||||
switch c.state {
|
switch c.state {
|
||||||
case StateReady:
|
case StateReady:
|
||||||
//yay!
|
//yay!
|
||||||
case StateStopped, StateStopping:
|
case StateStopped, StateStopping:
|
||||||
return nil, ErrAlreadyStopped
|
return ErrAlreadyStopped
|
||||||
case StateStarted:
|
case StateStarted:
|
||||||
return nil, ErrAlreadyStarted
|
return ErrAlreadyStarted
|
||||||
default:
|
default:
|
||||||
return nil, ErrUnknownState
|
return ErrUnknownState
|
||||||
}
|
}
|
||||||
|
|
||||||
// Activate the interface
|
// Activate the interface
|
||||||
err := c.f.activate()
|
err := c.f.activate()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
// Cancel before Close so a caller returning from Wait always observes a dead Context
|
||||||
|
c.cancel()
|
||||||
|
_ = c.f.Close()
|
||||||
c.state = StateStopped
|
c.state = StateStopped
|
||||||
return nil, err
|
return 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.
|
||||||
@@ -114,13 +114,9 @@ func (c *Control) Start() (func() error, error) {
|
|||||||
c.f.triggerShutdown = c.Stop
|
c.f.triggerShutdown = c.Stop
|
||||||
|
|
||||||
// Start reading packets.
|
// Start reading packets.
|
||||||
out, err := c.f.run()
|
c.f.run()
|
||||||
if err != nil {
|
|
||||||
c.state = StateStopped
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
c.state = StateStarted
|
c.state = StateStarted
|
||||||
return out, nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) State() RunState {
|
func (c *Control) State() RunState {
|
||||||
@@ -133,10 +129,26 @@ 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 tears nebula down, closing all tunnels and releasing everything it holds.
|
||||||
|
// Use Wait to block until the shutdown has completed.
|
||||||
|
// A Control that has been stopped cannot be started again, Start will return ErrAlreadyStopped.
|
||||||
func (c *Control) Stop() {
|
func (c *Control) Stop() {
|
||||||
c.stateLock.Lock()
|
c.stateLock.Lock()
|
||||||
if c.state != StateStarted {
|
switch c.state {
|
||||||
|
case StateStarted:
|
||||||
|
// Fall through to the full teardown below
|
||||||
|
|
||||||
|
case StateReady:
|
||||||
|
// Never started
|
||||||
|
c.cancel()
|
||||||
|
c.state = StateStopped
|
||||||
|
if err := c.f.Close(); err != nil {
|
||||||
|
c.l.Error("Close interface failed", "error", err)
|
||||||
|
}
|
||||||
|
c.stateLock.Unlock()
|
||||||
|
return
|
||||||
|
|
||||||
|
default:
|
||||||
c.stateLock.Unlock()
|
c.stateLock.Unlock()
|
||||||
// We are stopping or stopped already
|
// We are stopping or stopped already
|
||||||
return
|
return
|
||||||
@@ -145,19 +157,26 @@ func (c *Control) Stop() {
|
|||||||
c.state = StateStopping
|
c.state = StateStopping
|
||||||
c.stateLock.Unlock()
|
c.stateLock.Unlock()
|
||||||
|
|
||||||
// Stop the handshakeManager (and other services), to prevent new tunnels from
|
// Closing tunnels can be slow with a large hostmap, don't hold the lock for it
|
||||||
// being created while we're shutting them all down.
|
|
||||||
c.cancel()
|
c.cancel()
|
||||||
|
|
||||||
c.CloseAllTunnels(false)
|
c.CloseAllTunnels(false)
|
||||||
|
|
||||||
|
c.stateLock.Lock()
|
||||||
|
c.state = StateStopped
|
||||||
if err := c.f.Close(); err != nil {
|
if err := c.f.Close(); err != nil {
|
||||||
c.l.Error("Close interface failed", "error", err)
|
c.l.Error("Close interface failed", "error", err)
|
||||||
}
|
}
|
||||||
c.stateLock.Lock()
|
|
||||||
c.state = StateStopped
|
|
||||||
c.stateLock.Unlock()
|
c.stateLock.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Wait blocks until nebula has fully stopped, either via Stop or an internal fatal error,
|
||||||
|
// and returns the first fatal packet reader error if there was one.
|
||||||
|
// It is safe to call from multiple goroutines and at any point in the lifecycle,
|
||||||
|
// but a Wait on a Control that is never started and never stopped will block forever.
|
||||||
|
func (c *Control) Wait() error {
|
||||||
|
return c.f.wait()
|
||||||
|
}
|
||||||
|
|
||||||
// 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
|
||||||
func (c *Control) ShutdownBlock() {
|
func (c *Control) ShutdownBlock() {
|
||||||
sigChan := make(chan os.Signal, 1)
|
sigChan := make(chan os.Signal, 1)
|
||||||
@@ -170,8 +189,15 @@ func (c *Control) ShutdownBlock() {
|
|||||||
c.Stop()
|
c.Stop()
|
||||||
}
|
}
|
||||||
|
|
||||||
// RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change
|
// RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change.
|
||||||
func (c *Control) RebindUDPServer() {
|
func (c *Control) RebindUDPServer() {
|
||||||
|
c.stateLock.Lock()
|
||||||
|
defer c.stateLock.Unlock()
|
||||||
|
|
||||||
|
if c.state != StateStarted {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
_ = c.f.outside.Rebind()
|
_ = c.f.outside.Rebind()
|
||||||
|
|
||||||
// Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0
|
// Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0
|
||||||
|
|||||||
@@ -0,0 +1,292 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"net/netip"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gaissmai/bart"
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/routing"
|
||||||
|
"github.com/slackhq/nebula/test"
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
type fakeDevice struct {
|
||||||
|
closeOnce sync.Once
|
||||||
|
closedCh chan struct{}
|
||||||
|
closed bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func newFakeDevice() *fakeDevice {
|
||||||
|
return &fakeDevice{closedCh: make(chan struct{})}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read blocks until Close like a real tun with no traffic, then reports EOF
|
||||||
|
// the same way a closed device does
|
||||||
|
func (d *fakeDevice) Read(p []byte) (int, error) {
|
||||||
|
<-d.closedCh
|
||||||
|
return 0, io.EOF
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *fakeDevice) Write(p []byte) (int, error) { return len(p), nil }
|
||||||
|
|
||||||
|
func (d *fakeDevice) Close() error {
|
||||||
|
d.closeOnce.Do(func() {
|
||||||
|
d.closed = true
|
||||||
|
close(d.closedCh)
|
||||||
|
})
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *fakeDevice) Activate() error { return nil }
|
||||||
|
func (d *fakeDevice) Networks() []netip.Prefix { return nil }
|
||||||
|
func (d *fakeDevice) Name() string { return "fake" }
|
||||||
|
func (d *fakeDevice) RoutesFor(netip.Addr) routing.Gateways { return nil }
|
||||||
|
func (d *fakeDevice) SupportsMultiqueue() bool { return false }
|
||||||
|
func (d *fakeDevice) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
|
return nil, errors.New("unsupported")
|
||||||
|
}
|
||||||
|
|
||||||
|
// newReadyControl hand-builds the minimum Control that Main would have
|
||||||
|
// produced right before Start, including the construction token NewInterface
|
||||||
|
// takes so waiters block until Close releases the resources
|
||||||
|
func newReadyControl(t *testing.T) (*Control, *fakeDevice, *fakeConn) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
dev := newFakeDevice()
|
||||||
|
conn := &fakeConn{}
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
|
||||||
|
myVpnNet := netip.MustParsePrefix("10.128.0.1/16")
|
||||||
|
nt := new(bart.Lite)
|
||||||
|
nt.Insert(myVpnNet)
|
||||||
|
cs := &CertState{
|
||||||
|
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||||
|
myVpnNetworksTable: nt,
|
||||||
|
}
|
||||||
|
lh, err := NewLightHouseFromConfig(ctx, l, config.NewC(l), cs, nil, nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
f := &Interface{
|
||||||
|
ctx: ctx,
|
||||||
|
inside: dev,
|
||||||
|
outside: conn,
|
||||||
|
writers: []udp.Conn{conn},
|
||||||
|
readers: make([]io.ReadWriteCloser, 1),
|
||||||
|
routines: 1,
|
||||||
|
hostMap: newHostMap(l),
|
||||||
|
lightHouse: lh,
|
||||||
|
l: l,
|
||||||
|
}
|
||||||
|
f.wg.Add(1)
|
||||||
|
|
||||||
|
return &Control{
|
||||||
|
state: StateReady,
|
||||||
|
f: f,
|
||||||
|
l: l,
|
||||||
|
ctx: ctx,
|
||||||
|
cancel: cancel,
|
||||||
|
}, dev, conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_StopBeforeStart(t *testing.T) {
|
||||||
|
c, dev, conn := newReadyControl(t)
|
||||||
|
|
||||||
|
// A Stop on a never started control must release everything Main acquired
|
||||||
|
c.Stop()
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
assert.True(t, dev.closed, "the tun device should have been closed")
|
||||||
|
assert.True(t, conn.closed, "the udp socket should have been closed")
|
||||||
|
require.ErrorIs(t, c.ctx.Err(), context.Canceled, "the service context should have been cancelled")
|
||||||
|
|
||||||
|
// Wait must return promptly now that the resources are released
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
|
||||||
|
// A stopped control can never be started
|
||||||
|
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
||||||
|
|
||||||
|
// A second Stop is a harmless no-op
|
||||||
|
c.Stop()
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_WaitBlocksUntilStop(t *testing.T) {
|
||||||
|
c, _, _ := newReadyControl(t)
|
||||||
|
|
||||||
|
done := make(chan error, 1)
|
||||||
|
go func() { done <- c.Wait() }()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
t.Fatal("Wait returned before Stop")
|
||||||
|
case <-time.After(50 * time.Millisecond):
|
||||||
|
}
|
||||||
|
|
||||||
|
c.Stop()
|
||||||
|
select {
|
||||||
|
case err := <-done:
|
||||||
|
require.NoError(t, err)
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("Wait did not return after Stop")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type fakeConn struct {
|
||||||
|
closed bool
|
||||||
|
rebinds int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakeConn) Rebind() error { c.rebinds++; return nil }
|
||||||
|
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
||||||
|
func (c *fakeConn) ListenOut(_ udp.EncReader) error { return nil }
|
||||||
|
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
|
||||||
|
func (c *fakeConn) ReloadConfig(_ *config.C) {}
|
||||||
|
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
|
||||||
|
func (c *fakeConn) Close() error { c.closed = true; return nil }
|
||||||
|
|
||||||
|
type multiqueueDevice struct {
|
||||||
|
*fakeDevice
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *multiqueueDevice) SupportsMultiqueue() bool { return true }
|
||||||
|
|
||||||
|
func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
||||||
|
dev := &multiqueueDevice{fakeDevice: newFakeDevice()}
|
||||||
|
conn := &fakeConn{}
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
f := &Interface{
|
||||||
|
ctx: ctx,
|
||||||
|
inside: dev,
|
||||||
|
outside: conn,
|
||||||
|
writers: []udp.Conn{conn},
|
||||||
|
readers: make([]io.ReadWriteCloser, 2),
|
||||||
|
routines: 2,
|
||||||
|
l: test.NewLogger(),
|
||||||
|
}
|
||||||
|
f.wg.Add(1)
|
||||||
|
|
||||||
|
c := &Control{
|
||||||
|
state: StateReady,
|
||||||
|
f: f,
|
||||||
|
l: test.NewLogger(),
|
||||||
|
ctx: ctx,
|
||||||
|
cancel: cancel,
|
||||||
|
}
|
||||||
|
|
||||||
|
// The second reader fails to open, everything must be released
|
||||||
|
require.Error(t, c.Start())
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
assert.True(t, dev.closed, "the tun device should have been closed")
|
||||||
|
assert.True(t, conn.closed, "the udp socket should have been closed")
|
||||||
|
require.ErrorIs(t, c.ctx.Err(), context.Canceled)
|
||||||
|
|
||||||
|
// And Wait must not hang on the construction token
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInterface_CloseIsIdempotent(t *testing.T) {
|
||||||
|
dev := newFakeDevice()
|
||||||
|
f := &Interface{
|
||||||
|
inside: dev,
|
||||||
|
l: test.NewLogger(),
|
||||||
|
}
|
||||||
|
f.wg.Add(1)
|
||||||
|
|
||||||
|
require.NoError(t, f.Close())
|
||||||
|
assert.True(t, dev.closed)
|
||||||
|
|
||||||
|
// A second Close must not double release the wg token or the device
|
||||||
|
require.NoError(t, f.Close())
|
||||||
|
require.NoError(t, f.wait())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_FatalErrorReportsThroughWait(t *testing.T) {
|
||||||
|
c, dev, conn := newReadyControl(t)
|
||||||
|
|
||||||
|
// Mirror what Start wires up, without needing real packet readers
|
||||||
|
c.f.triggerShutdown = c.Stop
|
||||||
|
c.state = StateStarted
|
||||||
|
|
||||||
|
boom := errors.New("boom")
|
||||||
|
c.f.onFatal(boom)
|
||||||
|
|
||||||
|
require.ErrorIs(t, c.Wait(), boom)
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
assert.True(t, dev.closed)
|
||||||
|
assert.True(t, conn.closed)
|
||||||
|
|
||||||
|
// A second fatal error must not fire the shutdown again or replace the first
|
||||||
|
c.f.onFatal(errors.New("later"))
|
||||||
|
require.ErrorIs(t, c.Wait(), boom)
|
||||||
|
|
||||||
|
// Wait stays factual, a Stop after the death does not mask the error
|
||||||
|
c.Stop()
|
||||||
|
require.ErrorIs(t, c.Wait(), boom)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_ConcurrentStopAndStart(t *testing.T) {
|
||||||
|
c, _, _ := newReadyControl(t)
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for i := 0; i < 2; i++ {
|
||||||
|
wg.Go(func() { c.Stop() })
|
||||||
|
}
|
||||||
|
wg.Go(func() { _ = c.Start() })
|
||||||
|
wg.Go(func() {
|
||||||
|
_ = c.Wait()
|
||||||
|
// A returned Wait must always observe the final state, no matter how
|
||||||
|
// the race resolved
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
})
|
||||||
|
wg.Wait()
|
||||||
|
|
||||||
|
// However the race resolves, the control must end fully stopped with no
|
||||||
|
// panic and Wait must observe the final state
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_StartStopLifecycle(t *testing.T) {
|
||||||
|
c, dev, conn := newReadyControl(t)
|
||||||
|
|
||||||
|
require.NoError(t, c.Start())
|
||||||
|
assert.Equal(t, StateStarted, c.State())
|
||||||
|
require.ErrorIs(t, c.Start(), ErrAlreadyStarted)
|
||||||
|
|
||||||
|
// Stop must unpark the reader blocked in the device and release everything
|
||||||
|
c.Stop()
|
||||||
|
assert.Equal(t, StateStopped, c.State())
|
||||||
|
assert.True(t, dev.closed, "the tun device should have been closed")
|
||||||
|
assert.True(t, conn.closed, "the udp socket should have been closed")
|
||||||
|
require.ErrorIs(t, c.ctx.Err(), context.Canceled)
|
||||||
|
|
||||||
|
// The reader drained off a closed device, that is not a fatal error
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestControl_RebindIsGatedByState(t *testing.T) {
|
||||||
|
c, _, conn := newReadyControl(t)
|
||||||
|
|
||||||
|
// A rebind before Start reaches nothing, the interface is not up
|
||||||
|
c.RebindUDPServer()
|
||||||
|
assert.Equal(t, 0, conn.rebinds, "rebind before start must be a no-op")
|
||||||
|
|
||||||
|
require.NoError(t, c.Start())
|
||||||
|
c.RebindUDPServer()
|
||||||
|
assert.Equal(t, 1, conn.rebinds, "rebind while started must reach the conn")
|
||||||
|
|
||||||
|
// A rebind racing a completed stop must not touch the closed conn
|
||||||
|
c.Stop()
|
||||||
|
require.NoError(t, c.Wait())
|
||||||
|
c.RebindUDPServer()
|
||||||
|
assert.Equal(t, 1, conn.rebinds, "rebind after stop must be a no-op")
|
||||||
|
}
|
||||||
@@ -125,6 +125,14 @@ func (c *Control) GetHostmap() *HostMap {
|
|||||||
return c.f.hostMap
|
return c.f.hostMap
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetHostmapIndexCount returns the number of entries in the main hostmap Indexes table, holding
|
||||||
|
// the hostmap read lock so tests can poll it while connection manager churns tunnels.
|
||||||
|
func (c *Control) GetHostmapIndexCount() int {
|
||||||
|
c.f.hostMap.RLock()
|
||||||
|
defer c.f.hostMap.RUnlock()
|
||||||
|
return len(c.f.hostMap.Indexes)
|
||||||
|
}
|
||||||
|
|
||||||
func (c *Control) GetF() *Interface {
|
func (c *Control) GetF() *Interface {
|
||||||
return c.f
|
return c.f
|
||||||
}
|
}
|
||||||
|
|||||||
+34
-30
@@ -405,7 +405,7 @@ func TestStage1Race(t *testing.T) {
|
|||||||
|
|
||||||
r.Log("Spin until connection manager tears down a tunnel")
|
r.Log("Spin until connection manager tears down a tunnel")
|
||||||
|
|
||||||
for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 {
|
for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 {
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
t.Log("Connection manager hasn't ticked yet")
|
t.Log("Connection manager hasn't ticked yet")
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
@@ -453,9 +453,11 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
|||||||
|
|
||||||
r.Log("Nuke my hostmap")
|
r.Log("Nuke my hostmap")
|
||||||
myHostmap := myControl.GetHostmap()
|
myHostmap := myControl.GetHostmap()
|
||||||
|
myHostmap.Lock()
|
||||||
myHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
|
myHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
|
||||||
myHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
myHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
||||||
myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
||||||
|
myHostmap.Unlock()
|
||||||
|
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again")))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again")))
|
||||||
p = r.RouteForAllUntilTxTun(theirControl)
|
p = r.RouteForAllUntilTxTun(theirControl)
|
||||||
@@ -465,10 +467,10 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
|||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
r.Log("Wait for the dead index to go away")
|
r.Log("Wait for the dead index to go away")
|
||||||
start := len(theirControl.GetHostmap().Indexes)
|
start := theirControl.GetHostmapIndexCount()
|
||||||
for {
|
for {
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
if len(theirControl.GetHostmap().Indexes) < start {
|
if theirControl.GetHostmapIndexCount() < start {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
@@ -504,9 +506,11 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
|
|
||||||
r.Log("Nuke my hostmap")
|
r.Log("Nuke my hostmap")
|
||||||
theirHostmap := theirControl.GetHostmap()
|
theirHostmap := theirControl.GetHostmap()
|
||||||
|
theirHostmap.Lock()
|
||||||
theirHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
|
theirHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
|
||||||
theirHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
theirHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
||||||
theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
||||||
|
theirHostmap.Unlock()
|
||||||
|
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again")))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again")))
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
@@ -517,10 +521,10 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
r.Log("Wait for the dead index to go away")
|
r.Log("Wait for the dead index to go away")
|
||||||
start := len(myControl.GetHostmap().Indexes)
|
start := myControl.GetHostmapIndexCount()
|
||||||
for {
|
for {
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
if len(myControl.GetHostmap().Indexes) < start {
|
if myControl.GetHostmapIndexCount() < start {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
@@ -628,10 +632,10 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
r.Log("Close the tunnel")
|
r.Log("Close the tunnel")
|
||||||
relayControl.CloseTunnel(theirVpnIpNet[0].Addr(), true)
|
relayControl.CloseTunnel(theirVpnIpNet[0].Addr(), true)
|
||||||
|
|
||||||
start := len(myControl.GetHostmap().Indexes)
|
start := myControl.GetHostmapIndexCount()
|
||||||
curIndexes := len(myControl.GetHostmap().Indexes)
|
curIndexes := myControl.GetHostmapIndexCount()
|
||||||
for curIndexes >= start {
|
for curIndexes >= start {
|
||||||
curIndexes = len(myControl.GetHostmap().Indexes)
|
curIndexes = myControl.GetHostmapIndexCount()
|
||||||
r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes)
|
r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes)
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail")))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail")))
|
||||||
|
|
||||||
@@ -819,18 +823,18 @@ func TestStage1RaceRelays2(t *testing.T) {
|
|||||||
|
|
||||||
t.Log("Wait until we remove extra tunnels")
|
t.Log("Wait until we remove extra tunnels")
|
||||||
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
||||||
len(myControl.GetHostmap().Indexes),
|
myControl.GetHostmapIndexCount(),
|
||||||
len(theirControl.GetHostmap().Indexes),
|
theirControl.GetHostmapIndexCount(),
|
||||||
len(relayControl.GetHostmap().Indexes),
|
relayControl.GetHostmapIndexCount(),
|
||||||
)
|
)
|
||||||
hostInfos := len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes)
|
hostInfos := myControl.GetHostmapIndexCount() + theirControl.GetHostmapIndexCount() + relayControl.GetHostmapIndexCount()
|
||||||
retries := 60
|
retries := 60
|
||||||
for hostInfos > 6 && retries > 0 {
|
for hostInfos > 6 && retries > 0 {
|
||||||
hostInfos = len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes)
|
hostInfos = myControl.GetHostmapIndexCount() + theirControl.GetHostmapIndexCount() + relayControl.GetHostmapIndexCount()
|
||||||
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
||||||
len(myControl.GetHostmap().Indexes),
|
myControl.GetHostmapIndexCount(),
|
||||||
len(theirControl.GetHostmap().Indexes),
|
theirControl.GetHostmapIndexCount(),
|
||||||
len(relayControl.GetHostmap().Indexes),
|
relayControl.GetHostmapIndexCount(),
|
||||||
)
|
)
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
t.Log("Connection manager hasn't ticked yet")
|
t.Log("Connection manager hasn't ticked yet")
|
||||||
@@ -924,24 +928,24 @@ func TestRehandshakingRelays(t *testing.T) {
|
|||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl)
|
r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl)
|
||||||
// We should have two hostinfos on all sides
|
// We should have two hostinfos on all sides
|
||||||
for len(myControl.GetHostmap().Indexes) != 2 {
|
for myControl.GetHostmapIndexCount() != 2 {
|
||||||
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(myControl.GetHostmap().Indexes))
|
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", myControl.GetHostmapIndexCount())
|
||||||
r.Log("Assert the relay tunnel still works")
|
r.Log("Assert the relay tunnel still works")
|
||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.Log("yupitdoes")
|
r.Log("yupitdoes")
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
}
|
}
|
||||||
t.Logf("myControl hostinfos got cleaned up!")
|
t.Logf("myControl hostinfos got cleaned up!")
|
||||||
for len(theirControl.GetHostmap().Indexes) != 2 {
|
for theirControl.GetHostmapIndexCount() != 2 {
|
||||||
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(theirControl.GetHostmap().Indexes))
|
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", theirControl.GetHostmapIndexCount())
|
||||||
r.Log("Assert the relay tunnel still works")
|
r.Log("Assert the relay tunnel still works")
|
||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.Log("yupitdoes")
|
r.Log("yupitdoes")
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
}
|
}
|
||||||
t.Logf("theirControl hostinfos got cleaned up!")
|
t.Logf("theirControl hostinfos got cleaned up!")
|
||||||
for len(relayControl.GetHostmap().Indexes) != 2 {
|
for relayControl.GetHostmapIndexCount() != 2 {
|
||||||
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(relayControl.GetHostmap().Indexes))
|
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", relayControl.GetHostmapIndexCount())
|
||||||
r.Log("Assert the relay tunnel still works")
|
r.Log("Assert the relay tunnel still works")
|
||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.Log("yupitdoes")
|
r.Log("yupitdoes")
|
||||||
@@ -1029,24 +1033,24 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
|
|||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl)
|
r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl)
|
||||||
// We should have two hostinfos on all sides
|
// We should have two hostinfos on all sides
|
||||||
for len(myControl.GetHostmap().Indexes) != 2 {
|
for myControl.GetHostmapIndexCount() != 2 {
|
||||||
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(myControl.GetHostmap().Indexes))
|
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", myControl.GetHostmapIndexCount())
|
||||||
r.Log("Assert the relay tunnel still works")
|
r.Log("Assert the relay tunnel still works")
|
||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.Log("yupitdoes")
|
r.Log("yupitdoes")
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
}
|
}
|
||||||
t.Logf("myControl hostinfos got cleaned up!")
|
t.Logf("myControl hostinfos got cleaned up!")
|
||||||
for len(theirControl.GetHostmap().Indexes) != 2 {
|
for theirControl.GetHostmapIndexCount() != 2 {
|
||||||
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(theirControl.GetHostmap().Indexes))
|
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", theirControl.GetHostmapIndexCount())
|
||||||
r.Log("Assert the relay tunnel still works")
|
r.Log("Assert the relay tunnel still works")
|
||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.Log("yupitdoes")
|
r.Log("yupitdoes")
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
}
|
}
|
||||||
t.Logf("theirControl hostinfos got cleaned up!")
|
t.Logf("theirControl hostinfos got cleaned up!")
|
||||||
for len(relayControl.GetHostmap().Indexes) != 2 {
|
for relayControl.GetHostmapIndexCount() != 2 {
|
||||||
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(relayControl.GetHostmap().Indexes))
|
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", relayControl.GetHostmapIndexCount())
|
||||||
r.Log("Assert the relay tunnel still works")
|
r.Log("Assert the relay tunnel still works")
|
||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
r.Log("yupitdoes")
|
r.Log("yupitdoes")
|
||||||
@@ -1123,7 +1127,7 @@ func TestRehandshaking(t *testing.T) {
|
|||||||
theirConfig.ReloadConfigString(string(rc))
|
theirConfig.ReloadConfigString(string(rc))
|
||||||
|
|
||||||
r.Log("Spin until there is only 1 tunnel")
|
r.Log("Spin until there is only 1 tunnel")
|
||||||
for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 {
|
for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 {
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
t.Log("Connection manager hasn't ticked yet")
|
t.Log("Connection manager hasn't ticked yet")
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
@@ -1223,7 +1227,7 @@ func TestRehandshakingLoser(t *testing.T) {
|
|||||||
myConfig.ReloadConfigString(string(rc))
|
myConfig.ReloadConfigString(string(rc))
|
||||||
|
|
||||||
r.Log("Spin until there is only 1 tunnel")
|
r.Log("Spin until there is only 1 tunnel")
|
||||||
for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 {
|
for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 {
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
t.Log("Connection manager hasn't ticked yet")
|
t.Log("Connection manager hasn't ticked yet")
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
|
|||||||
+6
-6
@@ -43,8 +43,8 @@ func TestDropInactiveTunnels(t *testing.T) {
|
|||||||
r.Log("Go inactive and wait for the tunnels to get dropped")
|
r.Log("Go inactive and wait for the tunnels to get dropped")
|
||||||
waitStart := time.Now()
|
waitStart := time.Now()
|
||||||
for {
|
for {
|
||||||
myIndexes := len(myControl.GetHostmap().Indexes)
|
myIndexes := myControl.GetHostmapIndexCount()
|
||||||
theirIndexes := len(theirControl.GetHostmap().Indexes)
|
theirIndexes := theirControl.GetHostmapIndexCount()
|
||||||
if myIndexes == 0 && theirIndexes == 0 {
|
if myIndexes == 0 && theirIndexes == 0 {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
@@ -493,8 +493,8 @@ func TestCloseTunnelAuthenticated(t *testing.T) {
|
|||||||
|
|
||||||
waitStart := time.Now()
|
waitStart := time.Now()
|
||||||
for {
|
for {
|
||||||
myIndexes := len(myControl.GetHostmap().Indexes)
|
myIndexes := myControl.GetHostmapIndexCount()
|
||||||
theirIndexes := len(theirControl.GetHostmap().Indexes)
|
theirIndexes := theirControl.GetHostmapIndexCount()
|
||||||
if myIndexes == 0 && theirIndexes == 0 {
|
if myIndexes == 0 && theirIndexes == 0 {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
@@ -548,8 +548,8 @@ func TestCloseTunnelAuthenticated(t *testing.T) {
|
|||||||
r.Log("Injected bogus close tunnel. Let's see!")
|
r.Log("Injected bogus close tunnel. Let's see!")
|
||||||
waitStart = time.Now()
|
waitStart = time.Now()
|
||||||
for {
|
for {
|
||||||
myIndexes := len(myControl.GetHostmap().Indexes)
|
myIndexes := myControl.GetHostmapIndexCount()
|
||||||
theirIndexes := len(theirControl.GetHostmap().Indexes)
|
theirIndexes := theirControl.GetHostmapIndexCount()
|
||||||
if myIndexes == 0 {
|
if myIndexes == 0 {
|
||||||
t.Fatal("myIndexes should not be 0")
|
t.Fatal("myIndexes should not be 0")
|
||||||
}
|
}
|
||||||
|
|||||||
+9
-4
@@ -110,6 +110,15 @@ lighthouse:
|
|||||||
#- "1.1.1.1:4242"
|
#- "1.1.1.1:4242"
|
||||||
#- "1.2.3.4:0" # port will be replaced with the real listening port
|
#- "1.2.3.4:0" # port will be replaced with the real listening port
|
||||||
|
|
||||||
|
# Locally discovered addresses are checked against the MTU of the link they were found on. If the link cannot fit
|
||||||
|
# a full-size packet from the nebula tun device (`tun.mtu` plus encapsulation overhead, which is larger for relayed
|
||||||
|
# traffic) without fragmenting, a warning is logged.
|
||||||
|
# When omit_low_mtu_addrs is true, addresses whose links cannot fit normal nebula traffic are dropped from
|
||||||
|
# lighthouse reports entirely.
|
||||||
|
# Addresses that can fit normal nebula traffic but not relayed traffic are always still advertised.
|
||||||
|
# This does not apply to addresses listed in advertise_addrs.
|
||||||
|
#omit_low_mtu_addrs: false
|
||||||
|
|
||||||
# EXPERIMENTAL: This option may change or disappear in the future.
|
# EXPERIMENTAL: This option may change or disappear in the future.
|
||||||
# This setting allows us to "guess" what the remote might be for a host
|
# This setting allows us to "guess" what the remote might be for a host
|
||||||
# while we wait for the lighthouse response.
|
# while we wait for the lighthouse response.
|
||||||
@@ -242,10 +251,6 @@ tun:
|
|||||||
# When tun is disabled, a lighthouse can be started without a local tun interface (and therefore without root)
|
# When tun is disabled, a lighthouse can be started without a local tun interface (and therefore without root)
|
||||||
disabled: false
|
disabled: false
|
||||||
# Name of the device. If not set, a default will be chosen by the OS.
|
# Name of the device. If not set, a default will be chosen by the OS.
|
||||||
# For Linux: a single `%d` anywhere in the name is treated as a template and replaced with the
|
|
||||||
# lowest number that yields an unused device name (e.g. `nebula%d` becomes `nebula0`, then `nebula1`, and so on, `neb%dprod` becomes `neb0prod`).
|
|
||||||
# Only on Linux: `nebula%d` is the default if tun.dev is unset.
|
|
||||||
# The name, both before and after %d substitution, must be shorter than the kernel limit of 16 characters.
|
|
||||||
# For macOS: if set, must be in the form `utun[0-9]+`.
|
# For macOS: if set, must be in the form `utun[0-9]+`.
|
||||||
# For NetBSD: Required to be set, must be in the form `tun[0-9]+`
|
# For NetBSD: Required to be set, must be in the form `tun[0-9]+`
|
||||||
dev: nebula1
|
dev: nebula1
|
||||||
|
|||||||
@@ -430,14 +430,11 @@ func (hm *HandshakeManager) CheckAndComplete(hostinfo *HostInfo, handshakePacket
|
|||||||
// Check if we already have a tunnel with this vpn ip
|
// Check if we already have a tunnel with this vpn ip
|
||||||
existingHostInfo, found := hm.mainHostMap.Hosts[hostinfo.vpnAddrs[0]]
|
existingHostInfo, found := hm.mainHostMap.Hosts[hostinfo.vpnAddrs[0]]
|
||||||
if found && existingHostInfo != nil {
|
if found && existingHostInfo != nil {
|
||||||
testHostInfo := existingHostInfo
|
// Is it just a delayed handshake packet? Check every hostinfo we hold for this address.
|
||||||
for testHostInfo != nil {
|
for _, testHostInfo := range hm.mainHostMap.unlockedGetHostList(hostinfo.vpnAddrs[0]) {
|
||||||
// Is it just a delayed handshake packet?
|
|
||||||
if bytes.Equal(hostinfo.HandshakePacket[handshakePacket], testHostInfo.HandshakePacket[handshakePacket]) {
|
if bytes.Equal(hostinfo.HandshakePacket[handshakePacket], testHostInfo.HandshakePacket[handshakePacket]) {
|
||||||
return testHostInfo, ErrAlreadySeen
|
return testHostInfo, ErrAlreadySeen
|
||||||
}
|
}
|
||||||
|
|
||||||
testHostInfo = testHostInfo.next
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Is this a newer handshake?
|
// Is this a newer handshake?
|
||||||
|
|||||||
+163
-100
@@ -56,11 +56,20 @@ type Relay struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type HostMap struct {
|
type HostMap struct {
|
||||||
sync.RWMutex //Because we concurrently read and write to our maps
|
sync.RWMutex //Because we concurrently read and write to our maps
|
||||||
Indexes map[uint32]*HostInfo
|
Indexes map[uint32]*HostInfo
|
||||||
Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object
|
Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object
|
||||||
RemoteIndexes map[uint32]*HostInfo
|
RemoteIndexes map[uint32]*HostInfo
|
||||||
|
// Hosts maps a vpn address to its primary hostinfo, one entry per address we hold a tunnel
|
||||||
|
// for. moreHosts only has an entry while an address is held by 2 or more hostinfos and stores
|
||||||
|
// the full most-recent-first list; moreHosts[a][0] is always the same hostinfo as Hosts[a].
|
||||||
|
// Each address gets its own independent list, so a hostinfo owning multiple addresses can
|
||||||
|
// never corrupt another address's ordering the way the old shared next/prev chain could.
|
||||||
|
// Entries in moreHosts are only ever written by unlockedSetHostsForAddr; Hosts is written
|
||||||
|
// directly only in the single-hostinfo fast paths where moreHosts is known to have no entry,
|
||||||
|
// and unlockedDeleteHostInfo swaps either map for a fresh one when it fully drains.
|
||||||
Hosts map[netip.Addr]*HostInfo
|
Hosts map[netip.Addr]*HostInfo
|
||||||
|
moreHosts map[netip.Addr][]*HostInfo
|
||||||
preferredRanges atomic.Pointer[[]netip.Prefix]
|
preferredRanges atomic.Pointer[[]netip.Prefix]
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
@@ -266,10 +275,6 @@ type HostInfo struct {
|
|||||||
lastRoam time.Time
|
lastRoam time.Time
|
||||||
lastRoamRemote netip.AddrPort
|
lastRoamRemote netip.AddrPort
|
||||||
|
|
||||||
// Used to track other hostinfos for this vpn ip since only 1 can be primary
|
|
||||||
// Synchronised via hostmap lock and not the hostinfo lock.
|
|
||||||
next, prev *HostInfo
|
|
||||||
|
|
||||||
//TODO: in, out, and others might benefit from being an atomic.Int32. We could collapse connectionManager pendingDeletion, relayUsed, and in/out into this 1 thing
|
//TODO: in, out, and others might benefit from being an atomic.Int32. We could collapse connectionManager pendingDeletion, relayUsed, and in/out into this 1 thing
|
||||||
in, out, pendingDeletion atomic.Bool
|
in, out, pendingDeletion atomic.Bool
|
||||||
|
|
||||||
@@ -334,6 +339,7 @@ func newHostMap(l *slog.Logger) *HostMap {
|
|||||||
Relays: map[uint32]*HostInfo{},
|
Relays: map[uint32]*HostInfo{},
|
||||||
RemoteIndexes: map[uint32]*HostInfo{},
|
RemoteIndexes: map[uint32]*HostInfo{},
|
||||||
Hosts: map[netip.Addr]*HostInfo{},
|
Hosts: map[netip.Addr]*HostInfo{},
|
||||||
|
moreHosts: map[netip.Addr][]*HostInfo{},
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -382,13 +388,55 @@ func (hm *HostMap) EmitStats() {
|
|||||||
metrics.GetOrRegisterGauge("hostmap.main.relayIndexes", nil).Update(int64(relaysLen))
|
metrics.GetOrRegisterGauge("hostmap.main.relayIndexes", nil).Update(int64(relaysLen))
|
||||||
}
|
}
|
||||||
|
|
||||||
// DeleteHostInfo will fully unlink the hostinfo and return true if it was the final hostinfo for this vpn ip
|
// unlockedSetHostsForAddr stores the per-address hostinfo list (list[0] is the primary). An empty
|
||||||
|
// list removes the address. This is the one place Hosts and moreHosts are written together, keep
|
||||||
|
// it that way. Callers must hold the write lock.
|
||||||
|
func (hm *HostMap) unlockedSetHostsForAddr(addr netip.Addr, list []*HostInfo) {
|
||||||
|
if len(list) == 0 {
|
||||||
|
delete(hm.Hosts, addr)
|
||||||
|
delete(hm.moreHosts, addr)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
hm.Hosts[addr] = list[0]
|
||||||
|
if len(list) > 1 {
|
||||||
|
hm.moreHosts[addr] = list
|
||||||
|
} else {
|
||||||
|
delete(hm.moreHosts, addr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// unlockedGetHostList returns every hostinfo holding addr, primary first, or nil if we have no
|
||||||
|
// tunnel for addr. The common single-hostinfo case builds a fresh one element list, so keep this
|
||||||
|
// off the packet hot path; the primary is a direct Hosts read. Callers must hold the lock (read
|
||||||
|
// or write).
|
||||||
|
func (hm *HostMap) unlockedGetHostList(addr netip.Addr) []*HostInfo {
|
||||||
|
if list, ok := hm.moreHosts[addr]; ok {
|
||||||
|
return list
|
||||||
|
}
|
||||||
|
if h, ok := hm.Hosts[addr]; ok {
|
||||||
|
return []*HostInfo{h}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// removeHostInfo returns list with hi removed (order preserved), or list unchanged if hi is
|
||||||
|
// absent. It deletes in place: every mutator holds the hostmap write lock and no reader ever
|
||||||
|
// retains a slice across a mutation (readers iterate under RLock), so there is no snapshot to
|
||||||
|
// invalidate.
|
||||||
|
func removeHostInfo(list []*HostInfo, hi *HostInfo) []*HostInfo {
|
||||||
|
idx := slices.Index(list, hi)
|
||||||
|
if idx < 0 {
|
||||||
|
return list
|
||||||
|
}
|
||||||
|
return slices.Delete(list, idx, idx+1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteHostInfo will fully unlink the hostinfo and return true if no other hostinfo still holds
|
||||||
|
// any of its vpn addrs, meaning we no longer have a tunnel to the peer
|
||||||
func (hm *HostMap) DeleteHostInfo(hostinfo *HostInfo) bool {
|
func (hm *HostMap) DeleteHostInfo(hostinfo *HostInfo) bool {
|
||||||
// Delete the host itself, ensuring it's not modified anymore
|
// Delete the host itself, ensuring it's not modified anymore
|
||||||
hm.Lock()
|
hm.Lock()
|
||||||
// If we have a previous or next hostinfo then we are not the last one for this vpn ip
|
final := hm.unlockedDeleteHostInfo(hostinfo)
|
||||||
final := (hostinfo.next == nil && hostinfo.prev == nil)
|
|
||||||
hm.unlockedDeleteHostInfo(hostinfo)
|
|
||||||
hm.Unlock()
|
hm.Unlock()
|
||||||
|
|
||||||
return final
|
return final
|
||||||
@@ -400,71 +448,66 @@ func (hm *HostMap) MakePrimary(hostinfo *HostInfo) {
|
|||||||
hm.unlockedMakePrimary(hostinfo)
|
hm.unlockedMakePrimary(hostinfo)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) {
|
// unlockedMakePrimary reports whether hostinfo is (now) the primary for each of its addresses,
|
||||||
// Get the current primary, if it exists
|
// false only when it is no longer in the hostmap at all.
|
||||||
oldHostinfo := hm.Hosts[hostinfo.vpnAddrs[0]]
|
func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) bool {
|
||||||
|
// A hostinfo that is no longer in the hostmap must not be re-inserted here. Callers can race
|
||||||
// Every address in the hostinfo gets elevated to primary
|
// tunnel teardown, deciding to promote under the read lock and only taking the write lock
|
||||||
for _, vpnAddr := range hostinfo.vpnAddrs {
|
// after a delete fully unlinked the hostinfo (connection manager swapPrimary, AddRelay). Every
|
||||||
//NOTE: It is possible that we leave a dangling hostinfo here but connection manager works on
|
// live hostinfo is registered in Indexes by unlockedAddHostInfo, so this is a membership test.
|
||||||
// indexes so it should be fine.
|
if hm.Indexes[hostinfo.localIndexId] != hostinfo {
|
||||||
hm.Hosts[vpnAddr] = hostinfo
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
// If we are already primary then we won't bother re-linking
|
// Move hostinfo to the front (primary) of each of its address lists. The lists are
|
||||||
if oldHostinfo == hostinfo {
|
// independent per address, so this can never leave a dangling entry the way promoting
|
||||||
return
|
// against a single shared chain could.
|
||||||
}
|
|
||||||
|
|
||||||
// Unlink this hostinfo
|
|
||||||
if hostinfo.prev != nil {
|
|
||||||
hostinfo.prev.next = hostinfo.next
|
|
||||||
}
|
|
||||||
if hostinfo.next != nil {
|
|
||||||
hostinfo.next.prev = hostinfo.prev
|
|
||||||
}
|
|
||||||
|
|
||||||
// If there wasn't a previous primary then clear out any links
|
|
||||||
if oldHostinfo == nil {
|
|
||||||
hostinfo.next = nil
|
|
||||||
hostinfo.prev = nil
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Relink the hostinfo as primary
|
|
||||||
hostinfo.next = oldHostinfo
|
|
||||||
oldHostinfo.prev = hostinfo
|
|
||||||
hostinfo.prev = nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
|
||||||
isLastHostinfo := hostinfo.next == nil && hostinfo.prev == nil
|
|
||||||
|
|
||||||
for _, addr := range hostinfo.vpnAddrs {
|
for _, addr := range hostinfo.vpnAddrs {
|
||||||
if hm.Hosts[addr] != hostinfo {
|
if hm.Hosts[addr] == hostinfo {
|
||||||
|
// Already primary for this address, the list is already in the right order
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
if hostinfo.next != nil {
|
list := removeHostInfo(hm.unlockedGetHostList(addr), hostinfo)
|
||||||
// Promote the next hostinfo in the shared chain to primary for this address
|
list = append([]*HostInfo{hostinfo}, list...)
|
||||||
hm.Hosts[addr] = hostinfo.next
|
hm.unlockedSetHostsForAddr(addr, list)
|
||||||
} else {
|
}
|
||||||
delete(hm.Hosts, addr)
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// unlockedDeleteHostInfo removes hostinfo from every one of its address lists and from the index
|
||||||
|
// maps. It returns true if this was the last hostinfo for all of its addresses (we no longer have
|
||||||
|
// any tunnel to the peer), which the caller uses to decide whether to clear learned lighthouse
|
||||||
|
// state and disestablish relays.
|
||||||
|
func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
|
||||||
|
// Remove this hostinfo from each of its address lists. The lists are independent, so a
|
||||||
|
// sibling is never promoted to an address it does not own and no other list is touched.
|
||||||
|
final := true
|
||||||
|
for _, addr := range hostinfo.vpnAddrs {
|
||||||
|
if list, ok := hm.moreHosts[addr]; ok {
|
||||||
|
list = removeHostInfo(list, hostinfo)
|
||||||
|
hm.unlockedSetHostsForAddr(addr, list)
|
||||||
|
if len(list) > 0 {
|
||||||
|
final = false
|
||||||
|
}
|
||||||
|
} else if existing, ok := hm.Hosts[addr]; ok {
|
||||||
|
if existing == hostinfo {
|
||||||
|
// Common case, the only hostinfo for this address. moreHosts has no entry to clean up.
|
||||||
|
delete(hm.Hosts, addr)
|
||||||
|
} else {
|
||||||
|
// We don't hold this address but another hostinfo does, we still have a tunnel to the peer
|
||||||
|
final = false
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Go maps never shrink their buckets, replace fully drained maps so a node that churned
|
||||||
|
// through a large peer count gives the memory back. Same idiom as the index maps below.
|
||||||
if len(hm.Hosts) == 0 {
|
if len(hm.Hosts) == 0 {
|
||||||
hm.Hosts = map[netip.Addr]*HostInfo{}
|
hm.Hosts = map[netip.Addr]*HostInfo{}
|
||||||
}
|
}
|
||||||
|
if len(hm.moreHosts) == 0 {
|
||||||
// Splice this hostinfo out of the shared chain exactly once
|
hm.moreHosts = map[netip.Addr][]*HostInfo{}
|
||||||
if hostinfo.prev != nil {
|
|
||||||
hostinfo.prev.next = hostinfo.next
|
|
||||||
}
|
}
|
||||||
if hostinfo.next != nil {
|
|
||||||
hostinfo.next.prev = hostinfo.prev
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo.next = nil
|
|
||||||
hostinfo.prev = nil
|
|
||||||
|
|
||||||
// The remote index uses index ids outside our control so lets make sure we are only removing
|
// The remote index uses index ids outside our control so lets make sure we are only removing
|
||||||
// the remote index pointer here if it points to the hostinfo we are deleting
|
// the remote index pointer here if it points to the hostinfo we are deleting
|
||||||
@@ -488,7 +531,7 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
if isLastHostinfo {
|
if final {
|
||||||
// I have lost connectivity to my peers. My relay tunnel is likely broken. Mark the next
|
// I have lost connectivity to my peers. My relay tunnel is likely broken. Mark the next
|
||||||
// hops as 'Requested' so that new relay tunnels are created in the future.
|
// hops as 'Requested' so that new relay tunnels are created in the future.
|
||||||
hm.unlockedDisestablishVpnAddrRelayFor(hostinfo)
|
hm.unlockedDisestablishVpnAddrRelayFor(hostinfo)
|
||||||
@@ -497,6 +540,8 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
|||||||
for _, localRelayIdx := range hostinfo.relayState.CopyRelayForIdxs() {
|
for _, localRelayIdx := range hostinfo.relayState.CopyRelayForIdxs() {
|
||||||
delete(hm.Relays, localRelayIdx)
|
delete(hm.Relays, localRelayIdx)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return final
|
||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HostMap) QueryIndex(index uint32) *HostInfo {
|
func (hm *HostMap) QueryIndex(index uint32) *HostInfo {
|
||||||
@@ -540,19 +585,30 @@ func (hm *HostMap) QueryVpnAddrsRelayFor(targetIps []netip.Addr, relayHostIp net
|
|||||||
hm.RLock()
|
hm.RLock()
|
||||||
defer hm.RUnlock()
|
defer hm.RUnlock()
|
||||||
|
|
||||||
|
// This runs per relayed packet, so check the primary with a single map probe and only consult
|
||||||
|
// moreHosts when the primary can't relay for us.
|
||||||
h, ok := hm.Hosts[relayHostIp]
|
h, ok := hm.Hosts[relayHostIp]
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, nil, errors.New("unable to find host")
|
return nil, nil, errors.New("unable to find host")
|
||||||
}
|
}
|
||||||
|
|
||||||
for h != nil {
|
for _, targetIp := range targetIps {
|
||||||
for _, targetIp := range targetIps {
|
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
||||||
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
if ok && r.State == Established {
|
||||||
if ok && r.State == Established {
|
return h, r, nil
|
||||||
return h, r, nil
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if list, ok := hm.moreHosts[relayHostIp]; ok {
|
||||||
|
// list[0] is the primary we already checked
|
||||||
|
for _, h := range list[1:] {
|
||||||
|
for _, targetIp := range targetIps {
|
||||||
|
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
||||||
|
if ok && r.State == Established {
|
||||||
|
return h, r, nil
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
h = h.next
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil, nil, errors.New("unable to find host with relay")
|
return nil, nil, errors.New("unable to find host with relay")
|
||||||
@@ -560,20 +616,14 @@ func (hm *HostMap) QueryVpnAddrsRelayFor(targetIps []netip.Addr, relayHostIp net
|
|||||||
|
|
||||||
func (hm *HostMap) unlockedDisestablishVpnAddrRelayFor(hi *HostInfo) {
|
func (hm *HostMap) unlockedDisestablishVpnAddrRelayFor(hi *HostInfo) {
|
||||||
for _, relayHostIp := range hi.relayState.CopyRelayIps() {
|
for _, relayHostIp := range hi.relayState.CopyRelayIps() {
|
||||||
if h, ok := hm.Hosts[relayHostIp]; ok {
|
for _, h := range hm.unlockedGetHostList(relayHostIp) {
|
||||||
for h != nil {
|
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
||||||
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
|
||||||
h = h.next
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for _, rs := range hi.relayState.CopyAllRelayFor() {
|
for _, rs := range hi.relayState.CopyAllRelayFor() {
|
||||||
if rs.Type == ForwardingType {
|
if rs.Type == ForwardingType {
|
||||||
if h, ok := hm.Hosts[rs.PeerAddr]; ok {
|
for _, h := range hm.unlockedGetHostList(rs.PeerAddr) {
|
||||||
for h != nil {
|
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
||||||
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
|
||||||
h = h.next
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -623,22 +673,27 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HostMap) unlockedInnerAddHostInfo(vpnAddr netip.Addr, hostinfo *HostInfo, f *Interface) {
|
func (hm *HostMap) unlockedInnerAddHostInfo(vpnAddr netip.Addr, hostinfo *HostInfo, f *Interface) {
|
||||||
existing := hm.Hosts[vpnAddr]
|
existing, ok := hm.Hosts[vpnAddr]
|
||||||
hm.Hosts[vpnAddr] = hostinfo
|
if !ok {
|
||||||
|
// Common case, the first hostinfo for this address. moreHosts stays empty.
|
||||||
if existing != nil && existing != hostinfo {
|
hm.Hosts[vpnAddr] = hostinfo
|
||||||
hostinfo.next = existing
|
return
|
||||||
existing.prev = hostinfo
|
|
||||||
}
|
}
|
||||||
|
|
||||||
i := 1
|
// The new hostinfo becomes the primary for this address. Remove any stale copy of it first so
|
||||||
check := hostinfo
|
// we never hold a duplicate, then prepend.
|
||||||
for check != nil {
|
list, ok := hm.moreHosts[vpnAddr]
|
||||||
if i > MaxHostInfosPerVpnIp {
|
if !ok {
|
||||||
hm.unlockedDeleteHostInfo(check)
|
list = []*HostInfo{existing}
|
||||||
}
|
}
|
||||||
check = check.next
|
list = removeHostInfo(list, hostinfo)
|
||||||
i++
|
list = append([]*HostInfo{hostinfo}, list...)
|
||||||
|
hm.unlockedSetHostsForAddr(vpnAddr, list)
|
||||||
|
|
||||||
|
// Enforce the per-address cap by fully retiring the oldest hostinfo once we exceed it.
|
||||||
|
// Deleting it removes it from all of its addresses and the index maps, matching prior behavior.
|
||||||
|
if len(list) > MaxHostInfosPerVpnIp {
|
||||||
|
hm.unlockedDeleteHostInfo(list[len(list)-1])
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -814,9 +869,17 @@ func (i *HostInfo) logger(l *slog.Logger) *slog.Logger {
|
|||||||
|
|
||||||
// Utility functions
|
// Utility functions
|
||||||
|
|
||||||
func localAddrs(l *slog.Logger, allowList *LocalAllowList) []netip.Addr {
|
// localAddr is a locally discovered address candidate for lighthouse
|
||||||
|
// advertisement, along with details about the link it was found on.
|
||||||
|
type localAddr struct {
|
||||||
|
addr netip.Addr
|
||||||
|
ifName string
|
||||||
|
linkMTU int // MTU reported for the link, or <= 0 if unknown
|
||||||
|
}
|
||||||
|
|
||||||
|
func localAddrs(l *slog.Logger, allowList *LocalAllowList) []localAddr {
|
||||||
//FIXME: This function is pretty garbage
|
//FIXME: This function is pretty garbage
|
||||||
var finalAddrs []netip.Addr
|
var finalAddrs []localAddr
|
||||||
ifaces, _ := net.Interfaces()
|
ifaces, _ := net.Interfaces()
|
||||||
for _, i := range ifaces {
|
for _, i := range ifaces {
|
||||||
allow := allowList.AllowName(i.Name)
|
allow := allowList.AllowName(i.Name)
|
||||||
@@ -861,7 +924,7 @@ func localAddrs(l *slog.Logger, allowList *LocalAllowList) []netip.Addr {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
finalAddrs = append(finalAddrs, addr)
|
finalAddrs = append(finalAddrs, localAddr{addr: addr, ifName: i.Name, linkMTU: i.MTU})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+237
-181
@@ -2,6 +2,7 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"slices"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
@@ -10,78 +11,84 @@ import (
|
|||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// chainIds returns the localIndexIds of the hostinfos holding addr, primary (index 0) first. It
|
||||||
|
// also validates the Hosts/moreHosts sync contract on every call so a mutation that broke it
|
||||||
|
// fails fast.
|
||||||
|
func chainIds(t *testing.T, hm *HostMap, addr netip.Addr) []uint32 {
|
||||||
|
t.Helper()
|
||||||
|
assertHostMapInvariants(t, hm)
|
||||||
|
list := hm.unlockedGetHostList(addr)
|
||||||
|
ids := make([]uint32, len(list))
|
||||||
|
for i, h := range list {
|
||||||
|
ids[i] = h.localIndexId
|
||||||
|
}
|
||||||
|
return ids
|
||||||
|
}
|
||||||
|
|
||||||
|
// assertHostMapInvariants checks the Hosts/moreHosts contract: moreHosts only holds addresses
|
||||||
|
// with 2 or more hostinfos, its first entry is always the primary in Hosts, lists never hold
|
||||||
|
// duplicates, every hostinfo in a list owns the address and is registered in Indexes, and every
|
||||||
|
// indexed hostinfo is reachable through each of its addresses.
|
||||||
|
func assertHostMapInvariants(t *testing.T, hm *HostMap) {
|
||||||
|
t.Helper()
|
||||||
|
for addr, list := range hm.moreHosts {
|
||||||
|
require.GreaterOrEqualf(t, len(list), 2, "moreHosts[%s] must hold at least 2 hostinfos", addr)
|
||||||
|
require.Samef(t, hm.Hosts[addr], list[0], "moreHosts[%s][0] must match the primary in Hosts", addr)
|
||||||
|
seen := map[*HostInfo]bool{}
|
||||||
|
for _, h := range list {
|
||||||
|
require.NotNilf(t, h, "moreHosts[%s] must never hold a nil hostinfo", addr)
|
||||||
|
require.Falsef(t, seen[h], "moreHosts[%s] holds hostinfo %d twice", addr, h.localIndexId)
|
||||||
|
seen[h] = true
|
||||||
|
require.Samef(t, hm.Indexes[h.localIndexId], h, "moreHosts[%s] member %d is not registered in Indexes", addr, h.localIndexId)
|
||||||
|
require.Truef(t, slices.Contains(h.vpnAddrs, addr), "moreHosts[%s] member %d does not own the address", addr, h.localIndexId)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for addr, h := range hm.Hosts {
|
||||||
|
require.NotNilf(t, h, "Hosts[%s] must never be nil", addr)
|
||||||
|
require.Samef(t, hm.Indexes[h.localIndexId], h, "Hosts[%s] primary %d is not registered in Indexes", addr, h.localIndexId)
|
||||||
|
require.Truef(t, slices.Contains(h.vpnAddrs, addr), "Hosts[%s] primary (index %d) does not own the address", addr, h.localIndexId)
|
||||||
|
}
|
||||||
|
for idx, h := range hm.Indexes {
|
||||||
|
require.Equalf(t, idx, h.localIndexId, "Indexes[%d] holds hostinfo with localIndexId %d", idx, h.localIndexId)
|
||||||
|
for _, va := range h.vpnAddrs {
|
||||||
|
require.Truef(t, slices.Contains(hm.unlockedGetHostList(va), h), "indexed hostinfo %d is missing from the list for %s", idx, va)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestHostMap_MakePrimary(t *testing.T) {
|
func TestHostMap_MakePrimary(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
hm := newHostMap(l)
|
hm := newHostMap(l)
|
||||||
|
|
||||||
f := &Interface{}
|
f := &Interface{}
|
||||||
|
a := netip.MustParseAddr("0.0.0.1")
|
||||||
|
|
||||||
h1 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 1}
|
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
||||||
h2 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 2}
|
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2}
|
||||||
h3 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 3}
|
h3 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 3}
|
||||||
h4 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 4}
|
h4 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 4}
|
||||||
|
|
||||||
hm.unlockedAddHostInfo(h4, f)
|
hm.unlockedAddHostInfo(h4, f)
|
||||||
hm.unlockedAddHostInfo(h3, f)
|
hm.unlockedAddHostInfo(h3, f)
|
||||||
hm.unlockedAddHostInfo(h2, f)
|
hm.unlockedAddHostInfo(h2, f)
|
||||||
hm.unlockedAddHostInfo(h1, f)
|
hm.unlockedAddHostInfo(h1, f)
|
||||||
|
|
||||||
// Make sure we go h1 -> h2 -> h3 -> h4
|
// Most-recently-added is primary: h1, h2, h3, h4
|
||||||
prim := hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
assert.Equal(t, []uint32{1, 2, 3, 4}, chainIds(t, hm, a))
|
||||||
assert.Equal(t, h1.localIndexId, prim.localIndexId)
|
assert.Equal(t, h1, hm.QueryVpnAddr(a))
|
||||||
assert.Equal(t, h2.localIndexId, prim.next.localIndexId)
|
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h2.next.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h3.prev.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h3.next.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h4.prev.localIndexId)
|
|
||||||
assert.Nil(t, h4.next)
|
|
||||||
|
|
||||||
// Swap h3/middle to primary
|
// Swap the middle to primary: h3, h1, h2, h4
|
||||||
hm.MakePrimary(h3)
|
hm.MakePrimary(h3)
|
||||||
|
assert.Equal(t, []uint32{3, 1, 2, 4}, chainIds(t, hm, a))
|
||||||
|
assert.Equal(t, h3, hm.QueryVpnAddr(a))
|
||||||
|
|
||||||
// Make sure we go h3 -> h1 -> h2 -> h4
|
// Swap the tail to primary: h4, h3, h1, h2
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
|
||||||
assert.Equal(t, h3.localIndexId, prim.localIndexId)
|
|
||||||
assert.Equal(t, h1.localIndexId, prim.next.localIndexId)
|
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h2.localIndexId, h1.next.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h1.prev.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h2.next.localIndexId)
|
|
||||||
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h4.prev.localIndexId)
|
|
||||||
assert.Nil(t, h4.next)
|
|
||||||
|
|
||||||
// Swap h4/tail to primary
|
|
||||||
hm.MakePrimary(h4)
|
hm.MakePrimary(h4)
|
||||||
|
assert.Equal(t, []uint32{4, 3, 1, 2}, chainIds(t, hm, a))
|
||||||
|
|
||||||
// Make sure we go h4 -> h3 -> h1 -> h2
|
// Swapping the current primary again is a no-op
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
|
||||||
assert.Equal(t, h4.localIndexId, prim.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, prim.next.localIndexId)
|
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h1.localIndexId, h3.next.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h3.prev.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h1.next.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h1.prev.localIndexId)
|
|
||||||
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
|
|
||||||
assert.Nil(t, h2.next)
|
|
||||||
|
|
||||||
// Swap h4 again should be no-op
|
|
||||||
hm.MakePrimary(h4)
|
hm.MakePrimary(h4)
|
||||||
|
assert.Equal(t, []uint32{4, 3, 1, 2}, chainIds(t, hm, a))
|
||||||
// Make sure we go h4 -> h3 -> h1 -> h2
|
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
|
||||||
assert.Equal(t, h4.localIndexId, prim.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, prim.next.localIndexId)
|
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h1.localIndexId, h3.next.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h3.prev.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h1.next.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h1.prev.localIndexId)
|
|
||||||
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
|
|
||||||
assert.Nil(t, h2.next)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestHostMap_DeleteHostInfo(t *testing.T) {
|
func TestHostMap_DeleteHostInfo(t *testing.T) {
|
||||||
@@ -89,13 +96,14 @@ func TestHostMap_DeleteHostInfo(t *testing.T) {
|
|||||||
hm := newHostMap(l)
|
hm := newHostMap(l)
|
||||||
|
|
||||||
f := &Interface{}
|
f := &Interface{}
|
||||||
|
a := netip.MustParseAddr("0.0.0.1")
|
||||||
|
|
||||||
h1 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 1}
|
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
||||||
h2 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 2}
|
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2}
|
||||||
h3 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 3}
|
h3 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 3}
|
||||||
h4 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 4}
|
h4 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 4}
|
||||||
h5 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 5}
|
h5 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 5}
|
||||||
h6 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 6}
|
h6 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 6}
|
||||||
|
|
||||||
hm.unlockedAddHostInfo(h6, f)
|
hm.unlockedAddHostInfo(h6, f)
|
||||||
hm.unlockedAddHostInfo(h5, f)
|
hm.unlockedAddHostInfo(h5, f)
|
||||||
@@ -104,94 +112,110 @@ func TestHostMap_DeleteHostInfo(t *testing.T) {
|
|||||||
hm.unlockedAddHostInfo(h2, f)
|
hm.unlockedAddHostInfo(h2, f)
|
||||||
hm.unlockedAddHostInfo(h1, f)
|
hm.unlockedAddHostInfo(h1, f)
|
||||||
|
|
||||||
// h6 should be deleted
|
// h6 is evicted by the MaxHostInfosPerVpnIp cap; the rest are newest-first.
|
||||||
assert.Nil(t, h6.next)
|
assert.Nil(t, hm.QueryIndex(h6.localIndexId))
|
||||||
assert.Nil(t, h6.prev)
|
assert.Equal(t, []uint32{1, 2, 3, 4, 5}, chainIds(t, hm, a))
|
||||||
h := hm.QueryIndex(h6.localIndexId)
|
|
||||||
assert.Nil(t, h)
|
|
||||||
|
|
||||||
// Make sure we go h1 -> h2 -> h3 -> h4 -> h5
|
// Delete primary; not final since siblings remain.
|
||||||
prim := hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
assert.False(t, hm.DeleteHostInfo(h1))
|
||||||
assert.Equal(t, h1.localIndexId, prim.localIndexId)
|
assert.Equal(t, []uint32{2, 3, 4, 5}, chainIds(t, hm, a))
|
||||||
assert.Equal(t, h2.localIndexId, prim.next.localIndexId)
|
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h2.next.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h3.prev.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h3.next.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h4.prev.localIndexId)
|
|
||||||
assert.Equal(t, h5.localIndexId, h4.next.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
|
|
||||||
assert.Nil(t, h5.next)
|
|
||||||
|
|
||||||
// Delete primary
|
// Deleting the same hostinfo again must not report final while siblings remain and must not
|
||||||
hm.DeleteHostInfo(h1)
|
// disturb the list. The old chain code got this wrong: the first delete nil'd next/prev, so a
|
||||||
assert.Nil(t, h1.prev)
|
// second delete looked final and wiped lighthouse state out from under the live sibling.
|
||||||
assert.Nil(t, h1.next)
|
assert.False(t, hm.DeleteHostInfo(h1))
|
||||||
|
assert.Equal(t, []uint32{2, 3, 4, 5}, chainIds(t, hm, a))
|
||||||
|
|
||||||
// Make sure we go h2 -> h3 -> h4 -> h5
|
// Delete a middle node.
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
assert.False(t, hm.DeleteHostInfo(h3))
|
||||||
assert.Equal(t, h2.localIndexId, prim.localIndexId)
|
assert.Equal(t, []uint32{2, 4, 5}, chainIds(t, hm, a))
|
||||||
assert.Equal(t, h3.localIndexId, prim.next.localIndexId)
|
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h3.localIndexId, h2.next.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h3.prev.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h3.next.localIndexId)
|
|
||||||
assert.Equal(t, h3.localIndexId, h4.prev.localIndexId)
|
|
||||||
assert.Equal(t, h5.localIndexId, h4.next.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
|
|
||||||
assert.Nil(t, h5.next)
|
|
||||||
|
|
||||||
// Delete in the middle
|
// Delete the tail.
|
||||||
hm.DeleteHostInfo(h3)
|
assert.False(t, hm.DeleteHostInfo(h5))
|
||||||
assert.Nil(t, h3.prev)
|
assert.Equal(t, []uint32{2, 4}, chainIds(t, hm, a))
|
||||||
assert.Nil(t, h3.next)
|
|
||||||
|
|
||||||
// Make sure we go h2 -> h4 -> h5
|
// Delete the head; h4 remains and becomes primary.
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
assert.False(t, hm.DeleteHostInfo(h2))
|
||||||
assert.Equal(t, h2.localIndexId, prim.localIndexId)
|
assert.Equal(t, []uint32{4}, chainIds(t, hm, a))
|
||||||
assert.Equal(t, h4.localIndexId, prim.next.localIndexId)
|
assert.Equal(t, h4, hm.QueryVpnAddr(a))
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h4.localIndexId, h2.next.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h4.prev.localIndexId)
|
|
||||||
assert.Equal(t, h5.localIndexId, h4.next.localIndexId)
|
|
||||||
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
|
|
||||||
assert.Nil(t, h5.next)
|
|
||||||
|
|
||||||
// Delete the tail
|
// Delete the only remaining item; final is true and the address is gone.
|
||||||
hm.DeleteHostInfo(h5)
|
assert.True(t, hm.DeleteHostInfo(h4))
|
||||||
assert.Nil(t, h5.prev)
|
assert.Empty(t, chainIds(t, hm, a))
|
||||||
assert.Nil(t, h5.next)
|
assert.Nil(t, hm.QueryVpnAddr(a))
|
||||||
|
|
||||||
// Make sure we go h2 -> h4
|
// Deleting an already-gone hostinfo is still final; nothing holds the address anymore.
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
assert.True(t, hm.DeleteHostInfo(h4))
|
||||||
assert.Equal(t, h2.localIndexId, prim.localIndexId)
|
assert.Empty(t, chainIds(t, hm, a))
|
||||||
assert.Equal(t, h4.localIndexId, prim.next.localIndexId)
|
}
|
||||||
assert.Nil(t, prim.prev)
|
|
||||||
assert.Equal(t, h4.localIndexId, h2.next.localIndexId)
|
|
||||||
assert.Equal(t, h2.localIndexId, h4.prev.localIndexId)
|
|
||||||
assert.Nil(t, h4.next)
|
|
||||||
|
|
||||||
// Delete the head
|
// TestHostMap_MakePrimary_DeletedHostInfo covers promoting a hostinfo that lost a race with
|
||||||
hm.DeleteHostInfo(h2)
|
// tunnel teardown: swapPrimary and AddRelay decide to promote while holding a stale pointer and
|
||||||
assert.Nil(t, h2.prev)
|
// only take the write lock after a delete fully unlinked the hostinfo. MakePrimary must be a
|
||||||
assert.Nil(t, h2.next)
|
// no-op, not a resurrection that installs an unmanaged primary.
|
||||||
|
func TestHostMap_MakePrimary_DeletedHostInfo(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
hm := newHostMap(l)
|
||||||
|
f := &Interface{}
|
||||||
|
a := netip.MustParseAddr("0.0.0.1")
|
||||||
|
|
||||||
// Make sure we only have h4
|
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2}
|
||||||
assert.Equal(t, h4.localIndexId, prim.localIndexId)
|
hm.unlockedAddHostInfo(h1, f)
|
||||||
assert.Nil(t, prim.prev)
|
hm.unlockedAddHostInfo(h2, f)
|
||||||
assert.Nil(t, prim.next)
|
|
||||||
assert.Nil(t, h4.next)
|
|
||||||
|
|
||||||
// Delete the only item
|
// h1 is fully deleted while another goroutine still holds a pointer to it.
|
||||||
hm.DeleteHostInfo(h4)
|
assert.False(t, hm.DeleteHostInfo(h1))
|
||||||
assert.Nil(t, h4.prev)
|
assert.Equal(t, []uint32{2}, chainIds(t, hm, a))
|
||||||
assert.Nil(t, h4.next)
|
|
||||||
|
|
||||||
// Make sure we have nil
|
// The stale promote must not bring it back.
|
||||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
hm.MakePrimary(h1)
|
||||||
assert.Nil(t, prim)
|
assert.Equal(t, []uint32{2}, chainIds(t, hm, a))
|
||||||
|
assert.Equal(t, h2, hm.QueryVpnAddr(a))
|
||||||
|
assert.Nil(t, hm.QueryIndex(h1.localIndexId))
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestHostMap_QueryVpnAddrsRelayFor_NonPrimary makes sure a relay established on an older
|
||||||
|
// hostinfo is still found after a newer tunnel without relay state takes primary for the same
|
||||||
|
// address. The lookup checks the primary first and falls back to the rest of the list.
|
||||||
|
func TestHostMap_QueryVpnAddrsRelayFor_NonPrimary(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
hm := newHostMap(l)
|
||||||
|
f := &Interface{}
|
||||||
|
relayAddr := netip.MustParseAddr("0.0.0.9")
|
||||||
|
target := netip.MustParseAddr("0.0.0.1")
|
||||||
|
|
||||||
|
older := &HostInfo{
|
||||||
|
vpnAddrs: []netip.Addr{relayAddr},
|
||||||
|
localIndexId: 1,
|
||||||
|
relayState: RelayState{
|
||||||
|
relayForByAddr: map[netip.Addr]*Relay{},
|
||||||
|
relayForByIdx: map[uint32]*Relay{},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
older.relayState.InsertRelay(target, 100, &Relay{Type: ForwardingType, State: Established, LocalIndex: 100, PeerAddr: target})
|
||||||
|
hm.unlockedAddHostInfo(older, f)
|
||||||
|
|
||||||
|
// The relay is found on the primary.
|
||||||
|
h, r, err := hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, relayAddr)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, older, h)
|
||||||
|
assert.Equal(t, uint32(100), r.LocalIndex)
|
||||||
|
|
||||||
|
// A re-handshake with no relay state takes primary; the established relay on the older
|
||||||
|
// hostinfo must still be found through the fallback.
|
||||||
|
newer := &HostInfo{vpnAddrs: []netip.Addr{relayAddr}, localIndexId: 2}
|
||||||
|
hm.unlockedAddHostInfo(newer, f)
|
||||||
|
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, relayAddr))
|
||||||
|
|
||||||
|
h, r, err = hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, relayAddr)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, older, h)
|
||||||
|
assert.Equal(t, uint32(100), r.LocalIndex)
|
||||||
|
|
||||||
|
// No hostinfo at all is a plain miss.
|
||||||
|
_, _, err = hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, netip.MustParseAddr("0.0.0.42"))
|
||||||
|
require.Error(t, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestHostMap_DeleteHostInfo_MultipleVpnAddrs exercises the case where a hostinfo carries more than one
|
// TestHostMap_DeleteHostInfo_MultipleVpnAddrs exercises the case where a hostinfo carries more than one
|
||||||
@@ -216,32 +240,82 @@ func TestHostMap_DeleteHostInfo_MultipleVpnAddrs(t *testing.T) {
|
|||||||
hm.unlockedAddHostInfo(other, f)
|
hm.unlockedAddHostInfo(other, f)
|
||||||
hm.unlockedAddHostInfo(head, f)
|
hm.unlockedAddHostInfo(head, f)
|
||||||
|
|
||||||
// head is primary for both addresses, other is next in the shared chain
|
// head is primary for both addresses, other is next in each address's list.
|
||||||
assert.Equal(t, head.localIndexId, hm.QueryVpnAddr(a).localIndexId)
|
assert.Equal(t, head, hm.QueryVpnAddr(a))
|
||||||
assert.Equal(t, head.localIndexId, hm.QueryVpnAddr(b).localIndexId)
|
assert.Equal(t, head, hm.QueryVpnAddr(b))
|
||||||
assert.Equal(t, other.localIndexId, head.next.localIndexId)
|
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, a))
|
||||||
assert.Equal(t, head.localIndexId, other.prev.localIndexId)
|
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, b))
|
||||||
|
|
||||||
// Delete the head. other is still live, so it must become primary for BOTH addresses.
|
// Delete the head. other is still live, so it must become primary for BOTH addresses.
|
||||||
hm.DeleteHostInfo(head)
|
assert.False(t, hm.DeleteHostInfo(head))
|
||||||
|
assert.Equal(t, other, hm.QueryVpnAddr(a))
|
||||||
|
assert.Equal(t, other, hm.QueryVpnAddr(b))
|
||||||
|
assert.Equal(t, []uint32{1}, chainIds(t, hm, a))
|
||||||
|
assert.Equal(t, []uint32{1}, chainIds(t, hm, b))
|
||||||
|
|
||||||
// Pre-fix: QueryVpnAddr(b) came back nil here because the second address was deleted rather than
|
// head is fully removed from the index map.
|
||||||
// promoted, leaving other unreachable at b.
|
|
||||||
require.NotNil(t, hm.QueryVpnAddr(a))
|
|
||||||
require.NotNil(t, hm.QueryVpnAddr(b))
|
|
||||||
assert.Equal(t, other.localIndexId, hm.QueryVpnAddr(a).localIndexId)
|
|
||||||
assert.Equal(t, other.localIndexId, hm.QueryVpnAddr(b).localIndexId)
|
|
||||||
|
|
||||||
// other is now the only hostinfo in the chain
|
|
||||||
assert.Nil(t, other.prev)
|
|
||||||
assert.Nil(t, other.next)
|
|
||||||
|
|
||||||
// head is fully detached
|
|
||||||
assert.Nil(t, head.prev)
|
|
||||||
assert.Nil(t, head.next)
|
|
||||||
assert.Nil(t, hm.QueryIndex(head.localIndexId))
|
assert.Nil(t, hm.QueryIndex(head.localIndexId))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestHostMap_DeleteHostInfo_DivergentVpnAddrs covers chained hostinfos for the same peer whose
|
||||||
|
// vpnAddrs sets differ (a re-handshake cert added a second address). Deleting the superset node
|
||||||
|
// must not promote a sibling to an address it does not own.
|
||||||
|
func TestHostMap_DeleteHostInfo_DivergentVpnAddrs(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
hm := newHostMap(l)
|
||||||
|
f := &Interface{}
|
||||||
|
a := netip.MustParseAddr("0.0.0.1")
|
||||||
|
b := netip.MustParseAddr("0.0.0.2")
|
||||||
|
|
||||||
|
// sub owns only a; super (a newer handshake) owns a and b.
|
||||||
|
sub := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
||||||
|
super := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 2}
|
||||||
|
hm.unlockedAddHostInfo(sub, f)
|
||||||
|
hm.unlockedAddHostInfo(super, f)
|
||||||
|
|
||||||
|
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, a))
|
||||||
|
assert.Equal(t, []uint32{2}, chainIds(t, hm, b))
|
||||||
|
|
||||||
|
// Delete super: a promotes to sub (which owns it); b has no remaining owner and must be
|
||||||
|
// removed, not dangled at sub (which does not own b).
|
||||||
|
assert.False(t, hm.DeleteHostInfo(super))
|
||||||
|
assert.Equal(t, []uint32{1}, chainIds(t, hm, a))
|
||||||
|
assert.Empty(t, chainIds(t, hm, b))
|
||||||
|
assert.Equal(t, sub, hm.QueryVpnAddr(a))
|
||||||
|
assert.Nil(t, hm.QueryVpnAddr(b))
|
||||||
|
assert.Nil(t, hm.QueryIndex(super.localIndexId))
|
||||||
|
|
||||||
|
// Deleting sub cleans up fully.
|
||||||
|
assert.True(t, hm.DeleteHostInfo(sub))
|
||||||
|
assert.Nil(t, hm.QueryVpnAddr(a))
|
||||||
|
assertHostMapInvariants(t, hm)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestHostMap_AddDivergentOverlap covers a new hostinfo claiming addresses currently owned by two
|
||||||
|
// DIFFERENT hostinfos. The old single shared next/prev chain overwrote a pointer and orphaned one
|
||||||
|
// of them (in Indexes but unreachable via its address); independent per-address lists cannot.
|
||||||
|
func TestHostMap_AddDivergentOverlap(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
hm := newHostMap(l)
|
||||||
|
f := &Interface{}
|
||||||
|
a := netip.MustParseAddr("0.0.0.1")
|
||||||
|
b := netip.MustParseAddr("0.0.0.2")
|
||||||
|
|
||||||
|
hiA := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
||||||
|
hiP := &HostInfo{vpnAddrs: []netip.Addr{b}, localIndexId: 2}
|
||||||
|
hm.unlockedAddHostInfo(hiA, f)
|
||||||
|
hm.unlockedAddHostInfo(hiP, f)
|
||||||
|
|
||||||
|
hiB := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 3}
|
||||||
|
hm.unlockedAddHostInfo(hiB, f)
|
||||||
|
|
||||||
|
assert.Equal(t, []uint32{3, 1}, chainIds(t, hm, a))
|
||||||
|
assert.Equal(t, []uint32{3, 2}, chainIds(t, hm, b))
|
||||||
|
// hiA is still reachable via its address (not orphaned) and still indexed.
|
||||||
|
assert.Contains(t, chainIds(t, hm, a), hiA.localIndexId)
|
||||||
|
assert.NotNil(t, hm.QueryIndex(hiA.localIndexId))
|
||||||
|
}
|
||||||
|
|
||||||
// TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs verifies the MaxHostInfosPerVpnIp overflow prune
|
// TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs verifies the MaxHostInfosPerVpnIp overflow prune
|
||||||
// (unlockedInnerAddHostInfo calls unlockedDeleteHostInfo on the oldest node once the chain is too long)
|
// (unlockedInnerAddHostInfo calls unlockedDeleteHostInfo on the oldest node once the chain is too long)
|
||||||
// still behaves when hostinfos carry more than one vpnAddr. The pruned node is always the tail, so it is
|
// still behaves when hostinfos carry more than one vpnAddr. The pruned node is always the tail, so it is
|
||||||
@@ -267,32 +341,14 @@ func TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs(t *testing.T) {
|
|||||||
|
|
||||||
oldest := hostinfos[len(hostinfos)-1]
|
oldest := hostinfos[len(hostinfos)-1]
|
||||||
|
|
||||||
// The oldest hostinfo should have been pruned and fully detached
|
// The oldest hostinfo was pruned from both lists and the index map.
|
||||||
assert.Nil(t, oldest.next)
|
|
||||||
assert.Nil(t, oldest.prev)
|
|
||||||
assert.Nil(t, hm.QueryIndex(oldest.localIndexId))
|
assert.Nil(t, hm.QueryIndex(oldest.localIndexId))
|
||||||
|
|
||||||
// Both addresses resolve to the same head, and that head is one of the survivors (not the pruned one)
|
// Both addresses hold exactly MaxHostInfosPerVpnIp survivors in the same order; oldest is absent.
|
||||||
primA := hm.QueryVpnAddr(a)
|
require.Len(t, chainIds(t, hm, a), MaxHostInfosPerVpnIp)
|
||||||
primB := hm.QueryVpnAddr(b)
|
assert.Equal(t, chainIds(t, hm, a), chainIds(t, hm, b), "both addresses must list the same survivors in the same order")
|
||||||
require.NotNil(t, primA)
|
assert.NotContains(t, chainIds(t, hm, a), oldest.localIndexId)
|
||||||
require.NotNil(t, primB)
|
assert.Equal(t, hm.QueryVpnAddr(a), hm.QueryVpnAddr(b))
|
||||||
assert.Equal(t, primA.localIndexId, primB.localIndexId)
|
|
||||||
assert.NotEqual(t, oldest.localIndexId, primA.localIndexId)
|
|
||||||
|
|
||||||
// Walk the shared chain: exactly MaxHostInfosPerVpnIp survivors, no cycles, oldest absent
|
|
||||||
seen := map[uint32]struct{}{}
|
|
||||||
for h := primA; h != nil; h = h.next {
|
|
||||||
_, dup := seen[h.localIndexId]
|
|
||||||
require.False(t, dup, "cycle detected in hostinfo chain")
|
|
||||||
seen[h.localIndexId] = struct{}{}
|
|
||||||
if h.next != nil {
|
|
||||||
assert.Equal(t, h.localIndexId, h.next.prev.localIndexId, "prev pointer must mirror next")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
assert.Len(t, seen, MaxHostInfosPerVpnIp)
|
|
||||||
_, prunedStillPresent := seen[oldest.localIndexId]
|
|
||||||
assert.False(t, prunedStillPresent)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestHostMap_reload(t *testing.T) {
|
func TestHostMap_reload(t *testing.T) {
|
||||||
|
|||||||
+29
-14
@@ -215,6 +215,9 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
|
|
||||||
ifce.connectionManager.intf = ifce
|
ifce.connectionManager.intf = ifce
|
||||||
|
|
||||||
|
// Held until Close so waiting on the interface blocks until the resources are actually released
|
||||||
|
ifce.wg.Add(1)
|
||||||
|
|
||||||
return ifce, nil
|
return ifce, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -258,17 +261,16 @@ func (f *Interface) activate() error {
|
|||||||
f.readers[i] = reader
|
f.readers[i] = reader
|
||||||
}
|
}
|
||||||
|
|
||||||
f.wg.Add(1) // for us to wait on Close() to return
|
// On error the caller owns the cleanup, Control.Start cancels the service context
|
||||||
|
// before releasing our resources so a waiter never observes a live context
|
||||||
if err = f.inside.Activate(); err != nil {
|
if err = f.inside.Activate(); err != nil {
|
||||||
f.wg.Done()
|
|
||||||
f.inside.Close()
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) run() (func() error, 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() {
|
f.wg.Go(func() {
|
||||||
@@ -283,13 +285,14 @@ func (f *Interface) run() (func() error, error) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
return func() error {
|
}
|
||||||
f.wg.Wait()
|
|
||||||
if e := f.fatalErr.Load(); e != nil {
|
func (f *Interface) wait() error {
|
||||||
return *e
|
f.wg.Wait()
|
||||||
}
|
if e := f.fatalErr.Load(); e != nil {
|
||||||
return nil
|
return *e
|
||||||
}, nil
|
}
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// onFatal stores the first fatal reader error, and calls triggerShutdown if it was the first one
|
// onFatal stores the first fatal reader error, and calls triggerShutdown if it was the first one
|
||||||
@@ -322,7 +325,10 @@ func (f *Interface) listenOut(i int) {
|
|||||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
||||||
})
|
})
|
||||||
|
|
||||||
if err != nil && !f.closed.Load() {
|
// An error after teardown began is shutdown noise, the closed flag covers resources
|
||||||
|
// Close releases itself and the cancelled ctx covers ones torn down by their owners
|
||||||
|
// reacting to it, like the user device pipes
|
||||||
|
if err != nil && !f.closed.Load() && f.ctx.Err() == nil {
|
||||||
f.l.Error("Error while reading inbound packet, closing", "error", err)
|
f.l.Error("Error while reading inbound packet, closing", "error", err)
|
||||||
f.onFatal(err)
|
f.onFatal(err)
|
||||||
}
|
}
|
||||||
@@ -341,7 +347,8 @@ 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() {
|
// Same shutdown noise handling as listenOut
|
||||||
|
if !f.closed.Load() && f.ctx.Err() == nil {
|
||||||
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i)
|
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i)
|
||||||
f.onFatal(err)
|
f.onFatal(err)
|
||||||
}
|
}
|
||||||
@@ -542,9 +549,15 @@ func (f *Interface) GetCertState() *CertState {
|
|||||||
return f.pki.getCertState()
|
return f.pki.getCertState()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Close releases the interface's resources: the udp sockets and the tun device.
|
||||||
|
// It is idempotent and safe to call at any point in the lifecycle, including on an interface that never activated,
|
||||||
|
// calls after the first return nil without doing anything.
|
||||||
func (f *Interface) Close() error {
|
func (f *Interface) Close() error {
|
||||||
|
if !f.closed.CompareAndSwap(false, true) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
var errs []error
|
var errs []error
|
||||||
f.closed.Store(true)
|
|
||||||
|
|
||||||
// Release the udp readers
|
// Release the udp readers
|
||||||
for i, u := range f.writers {
|
for i, u := range f.writers {
|
||||||
@@ -560,6 +573,8 @@ func (f *Interface) Close() error {
|
|||||||
if closeErr != nil {
|
if closeErr != nil {
|
||||||
errs = append(errs, closeErr)
|
errs = append(errs, closeErr)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Release the construction token so waiters know the resources are gone
|
||||||
f.wg.Done()
|
f.wg.Done()
|
||||||
return errors.Join(errs...)
|
return errors.Join(errs...)
|
||||||
}
|
}
|
||||||
|
|||||||
+143
-4
@@ -19,8 +19,11 @@ import (
|
|||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/logging"
|
"github.com/slackhq/nebula/logging"
|
||||||
|
"github.com/slackhq/nebula/overlay"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
|
"golang.org/x/net/ipv4"
|
||||||
|
"golang.org/x/net/ipv6"
|
||||||
)
|
)
|
||||||
|
|
||||||
var ErrHostNotKnown = errors.New("host not known")
|
var ErrHostNotKnown = errors.New("host not known")
|
||||||
@@ -65,6 +68,18 @@ type LightHouse struct {
|
|||||||
|
|
||||||
advertiseAddrs atomic.Pointer[[]netip.AddrPort]
|
advertiseAddrs atomic.Pointer[[]netip.AddrPort]
|
||||||
|
|
||||||
|
// tunMTU mirrors tun.mtu so locally discovered addrs can be checked for
|
||||||
|
// links too small to carry a full-size nebula packet without fragmenting.
|
||||||
|
tunMTU atomic.Int64
|
||||||
|
// omitLowMTUAddrs drops such addrs from lighthouse updates (and demotes
|
||||||
|
// the associated warnings to debug logs) instead of advertising them.
|
||||||
|
omitLowMTUAddrs atomic.Bool
|
||||||
|
// mtuWarned tracks the last classification logged per local addr so a
|
||||||
|
// warning is only emitted when the classification changes, not on every
|
||||||
|
// periodic update.
|
||||||
|
mtuWarnLock sync.Mutex
|
||||||
|
mtuWarned map[mtuWarnKey]linkMTUTier
|
||||||
|
|
||||||
// Addr's of relays that can be used by peers to access me
|
// Addr's of relays that can be used by peers to access me
|
||||||
relaysForMe atomic.Pointer[[]netip.Addr]
|
relaysForMe atomic.Pointer[[]netip.Addr]
|
||||||
|
|
||||||
@@ -105,6 +120,7 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
|
|||||||
punchy: p,
|
punchy: p,
|
||||||
updateTrigger: make(chan struct{}, 1),
|
updateTrigger: make(chan struct{}, 1),
|
||||||
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
|
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
|
||||||
|
mtuWarned: make(map[mtuWarnKey]linkMTUTier),
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
lighthouses := make([]netip.Addr, 0)
|
lighthouses := make([]netip.Addr, 0)
|
||||||
@@ -216,6 +232,23 @@ func (lh *LightHouse) reload(c *config.C, initial bool) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if initial || c.HasChanged("tun.mtu") || c.HasChanged("lighthouse.omit_low_mtu_addrs") {
|
||||||
|
lh.tunMTU.Store(int64(c.GetInt("tun.mtu", overlay.DefaultMTU)))
|
||||||
|
lh.omitLowMTUAddrs.Store(c.GetBool("lighthouse.omit_low_mtu_addrs", false))
|
||||||
|
|
||||||
|
// Re-log any addrs whose links are still too small under the new values
|
||||||
|
lh.mtuWarnLock.Lock()
|
||||||
|
clear(lh.mtuWarned)
|
||||||
|
lh.mtuWarnLock.Unlock()
|
||||||
|
|
||||||
|
if !initial {
|
||||||
|
lh.l.Info("tun.mtu and/or lighthouse.omit_low_mtu_addrs has changed",
|
||||||
|
"tunMTU", lh.tunMTU.Load(),
|
||||||
|
"omitLowMTUAddrs", lh.omitLowMTUAddrs.Load(),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if initial || c.HasChanged("lighthouse.interval") {
|
if initial || c.HasChanged("lighthouse.interval") {
|
||||||
lh.interval.Store(int64(c.GetInt("lighthouse.interval", 10)))
|
lh.interval.Store(int64(c.GetInt("lighthouse.interval", 10)))
|
||||||
|
|
||||||
@@ -905,6 +938,108 @@ func (lh *LightHouse) TriggerUpdate() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// linkMTUTier classifies how well a local addr's link MTU can carry
|
||||||
|
// full-size nebula packets built from a tun packet of tun.mtu bytes.
|
||||||
|
type linkMTUTier uint8
|
||||||
|
|
||||||
|
// mtuWarnKey identifies a local addr for MTU warning dedup purposes. The
|
||||||
|
// interface name is included because the same addr can exist on multiple
|
||||||
|
// links with different MTUs.
|
||||||
|
type mtuWarnKey struct {
|
||||||
|
ifName string
|
||||||
|
addr netip.Addr
|
||||||
|
}
|
||||||
|
|
||||||
|
const (
|
||||||
|
// The link can carry both normal and relayed nebula traffic
|
||||||
|
linkMTUOk linkMTUTier = iota
|
||||||
|
// The link can carry normal nebula traffic, but relayed traffic (which
|
||||||
|
// adds a second layer of encapsulation) will not fit
|
||||||
|
linkMTUTooSmallForRelay
|
||||||
|
// Even normal nebula traffic will not fit
|
||||||
|
linkMTUTooSmall
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// Both AES-256-GCM and ChaCha20-Poly1305 append a 16 byte AEAD tag
|
||||||
|
cipherTagLen = 16
|
||||||
|
udpHeaderLen = 8
|
||||||
|
)
|
||||||
|
|
||||||
|
// requiredLinkMTU returns the minimum underlay link MTU that can carry a
|
||||||
|
// full-size tun packet to an addr of the given family without fragmentation,
|
||||||
|
// both directly and via a relay (which wraps the packet in a second nebula
|
||||||
|
// header and AEAD tag).
|
||||||
|
func requiredLinkMTU(tunMTU int, is4 bool) (direct, relayed int) {
|
||||||
|
ipHeaderLen := ipv6.HeaderLen
|
||||||
|
if is4 {
|
||||||
|
ipHeaderLen = ipv4.HeaderLen
|
||||||
|
}
|
||||||
|
|
||||||
|
direct = tunMTU + header.Len + cipherTagLen + udpHeaderLen + ipHeaderLen
|
||||||
|
relayed = direct + header.Len + cipherTagLen
|
||||||
|
return direct, relayed
|
||||||
|
}
|
||||||
|
|
||||||
|
// checkLocalLinkMTU classifies e's link MTU, logs when the classification
|
||||||
|
// changes, and reports whether e should be advertised to lighthouses.
|
||||||
|
func (lh *LightHouse) checkLocalLinkMTU(e localAddr) bool {
|
||||||
|
tunMTU := int(lh.tunMTU.Load())
|
||||||
|
omit := lh.omitLowMTUAddrs.Load()
|
||||||
|
|
||||||
|
tier := linkMTUOk
|
||||||
|
direct, relayed := requiredLinkMTU(tunMTU, e.addr.Is4())
|
||||||
|
if e.linkMTU > 0 { // links with an unknown MTU are advertised as-is
|
||||||
|
if e.linkMTU < direct {
|
||||||
|
tier = linkMTUTooSmall
|
||||||
|
} else if e.linkMTU < relayed {
|
||||||
|
tier = linkMTUTooSmallForRelay
|
||||||
|
}
|
||||||
|
}
|
||||||
|
advertise := tier != linkMTUTooSmall || !omit
|
||||||
|
|
||||||
|
key := mtuWarnKey{ifName: e.ifName, addr: e.addr}
|
||||||
|
lh.mtuWarnLock.Lock()
|
||||||
|
changed := lh.mtuWarned[key] != tier
|
||||||
|
if changed {
|
||||||
|
if tier == linkMTUOk {
|
||||||
|
delete(lh.mtuWarned, key)
|
||||||
|
} else {
|
||||||
|
lh.mtuWarned[key] = tier
|
||||||
|
}
|
||||||
|
}
|
||||||
|
lh.mtuWarnLock.Unlock()
|
||||||
|
|
||||||
|
if !changed || tier == linkMTUOk {
|
||||||
|
return advertise
|
||||||
|
}
|
||||||
|
|
||||||
|
level := slog.LevelWarn
|
||||||
|
if omit {
|
||||||
|
level = slog.LevelDebug
|
||||||
|
}
|
||||||
|
|
||||||
|
if lh.l.Enabled(context.Background(), level) {
|
||||||
|
msg := "Link MTU too small for nebula traffic, expect fragmentation or drops"
|
||||||
|
if !advertise {
|
||||||
|
msg = "Omitting addr with too-small link MTU from lighthouse report"
|
||||||
|
} else if tier == linkMTUTooSmallForRelay {
|
||||||
|
msg = "Link MTU too small for relayed nebula traffic"
|
||||||
|
}
|
||||||
|
|
||||||
|
lh.l.Log(context.Background(), level, msg,
|
||||||
|
"localAddr", e.addr,
|
||||||
|
"interface", e.ifName,
|
||||||
|
"linkMTU", e.linkMTU,
|
||||||
|
"requiredMTU", direct,
|
||||||
|
"requiredRelayMTU", relayed,
|
||||||
|
"tunMTU", tunMTU,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
return advertise
|
||||||
|
}
|
||||||
|
|
||||||
func (lh *LightHouse) SendUpdate() {
|
func (lh *LightHouse) SendUpdate() {
|
||||||
var v4 []*V4AddrPort
|
var v4 []*V4AddrPort
|
||||||
var v6 []*V6AddrPort
|
var v6 []*V6AddrPort
|
||||||
@@ -919,15 +1054,19 @@ func (lh *LightHouse) SendUpdate() {
|
|||||||
|
|
||||||
lal := lh.GetLocalAllowList()
|
lal := lh.GetLocalAllowList()
|
||||||
for _, e := range localAddrs(lh.l, lal) {
|
for _, e := range localAddrs(lh.l, lal) {
|
||||||
if lh.myVpnNetworksTable.Contains(e) {
|
if lh.myVpnNetworksTable.Contains(e.addr) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if !lh.checkLocalLinkMTU(e) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Only add addrs that aren't my VPN/tun networks
|
// Only add addrs that aren't my VPN/tun networks
|
||||||
if e.Is4() {
|
if e.addr.Is4() {
|
||||||
v4 = append(v4, netAddrToProtoV4AddrPort(e, uint16(lh.nebulaPort)))
|
v4 = append(v4, netAddrToProtoV4AddrPort(e.addr, uint16(lh.nebulaPort)))
|
||||||
} else {
|
} else {
|
||||||
v6 = append(v6, netAddrToProtoV6AddrPort(e, uint16(lh.nebulaPort)))
|
v6 = append(v6, netAddrToProtoV6AddrPort(e.addr, uint16(lh.nebulaPort)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -738,3 +738,86 @@ func TestLighthouse_DeletesWork(t *testing.T) {
|
|||||||
out = lh.Query(testHost)
|
out = lh.Query(testHost)
|
||||||
assert.Nil(t, out)
|
assert.Nil(t, out)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func Test_requiredLinkMTU(t *testing.T) {
|
||||||
|
// tun packet + nebula header (16) + AEAD tag (16) + udp (8) + ip header
|
||||||
|
direct, relayed := requiredLinkMTU(1300, true)
|
||||||
|
assert.Equal(t, 1360, direct)
|
||||||
|
assert.Equal(t, 1392, relayed)
|
||||||
|
|
||||||
|
direct, relayed = requiredLinkMTU(1300, false)
|
||||||
|
assert.Equal(t, 1380, direct)
|
||||||
|
assert.Equal(t, 1412, relayed)
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_checkLocalLinkMTU(t *testing.T) {
|
||||||
|
lh := &LightHouse{l: test.NewLogger(), mtuWarned: make(map[mtuWarnKey]linkMTUTier)}
|
||||||
|
lh.tunMTU.Store(1300)
|
||||||
|
|
||||||
|
v4 := netip.MustParseAddr("192.168.1.2")
|
||||||
|
v6 := netip.MustParseAddr("fd00::2")
|
||||||
|
mkAddr := func(a netip.Addr, mtu int) localAddr {
|
||||||
|
return localAddr{addr: a, ifName: "test0", linkMTU: mtu}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Plenty of room, no state recorded
|
||||||
|
assert.True(t, lh.checkLocalLinkMTU(mkAddr(v4, 1500)))
|
||||||
|
assert.Empty(t, lh.mtuWarned)
|
||||||
|
|
||||||
|
// Unknown link MTU is not classified
|
||||||
|
assert.True(t, lh.checkLocalLinkMTU(mkAddr(v4, 0)))
|
||||||
|
assert.Empty(t, lh.mtuWarned)
|
||||||
|
|
||||||
|
// Too small for even normal traffic, still advertised by default
|
||||||
|
assert.True(t, lh.checkLocalLinkMTU(mkAddr(v4, 1359)))
|
||||||
|
assert.Equal(t, linkMTUTooSmall, lh.mtuWarned[mtuWarnKey{ifName: "test0", addr: v4}])
|
||||||
|
|
||||||
|
// Fits normal traffic but not relayed traffic
|
||||||
|
assert.True(t, lh.checkLocalLinkMTU(mkAddr(v4, 1360)))
|
||||||
|
assert.Equal(t, linkMTUTooSmallForRelay, lh.mtuWarned[mtuWarnKey{ifName: "test0", addr: v4}])
|
||||||
|
|
||||||
|
// Exactly enough for relayed traffic clears the state
|
||||||
|
assert.True(t, lh.checkLocalLinkMTU(mkAddr(v4, 1392)))
|
||||||
|
assert.Empty(t, lh.mtuWarned)
|
||||||
|
|
||||||
|
// v6 addrs need 20 more bytes of headroom
|
||||||
|
assert.True(t, lh.checkLocalLinkMTU(mkAddr(v6, 1380)))
|
||||||
|
assert.Equal(t, linkMTUTooSmallForRelay, lh.mtuWarned[mtuWarnKey{ifName: "test0", addr: v6}])
|
||||||
|
assert.True(t, lh.checkLocalLinkMTU(mkAddr(v6, 1379)))
|
||||||
|
assert.Equal(t, linkMTUTooSmall, lh.mtuWarned[mtuWarnKey{ifName: "test0", addr: v6}])
|
||||||
|
assert.True(t, lh.checkLocalLinkMTU(mkAddr(v6, 1412)))
|
||||||
|
assert.Empty(t, lh.mtuWarned)
|
||||||
|
|
||||||
|
// With omit enabled, only addrs that can't fit normal traffic are dropped
|
||||||
|
lh.omitLowMTUAddrs.Store(true)
|
||||||
|
assert.False(t, lh.checkLocalLinkMTU(mkAddr(v4, 1359)))
|
||||||
|
assert.Equal(t, linkMTUTooSmall, lh.mtuWarned[mtuWarnKey{ifName: "test0", addr: v4}])
|
||||||
|
assert.True(t, lh.checkLocalLinkMTU(mkAddr(v4, 1360)))
|
||||||
|
assert.Equal(t, linkMTUTooSmallForRelay, lh.mtuWarned[mtuWarnKey{ifName: "test0", addr: v4}])
|
||||||
|
assert.False(t, lh.checkLocalLinkMTU(mkAddr(v4, 1359)))
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_lighthouseMTUConfig(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
myVpnNet := netip.MustParsePrefix("10.128.0.1/16")
|
||||||
|
nt := new(bart.Lite)
|
||||||
|
nt.Insert(myVpnNet)
|
||||||
|
cs := &CertState{
|
||||||
|
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||||
|
myVpnNetworksTable: nt,
|
||||||
|
}
|
||||||
|
|
||||||
|
c := config.NewC(l)
|
||||||
|
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, int64(1300), lh.tunMTU.Load())
|
||||||
|
assert.False(t, lh.omitLowMTUAddrs.Load())
|
||||||
|
|
||||||
|
c = config.NewC(l)
|
||||||
|
c.Settings["tun"] = map[string]any{"mtu": 8000}
|
||||||
|
c.Settings["lighthouse"] = map[string]any{"omit_low_mtu_addrs": true}
|
||||||
|
lh, err = NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, int64(8000), lh.tunMTU.Load())
|
||||||
|
assert.True(t, lh.omitLowMTUAddrs.Load())
|
||||||
|
}
|
||||||
|
|||||||
@@ -130,6 +130,17 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
udpConns := make([]udp.Conn, routines)
|
udpConns := make([]udp.Conn, routines)
|
||||||
port := c.GetInt("listen.port", 0)
|
port := c.GetInt("listen.port", 0)
|
||||||
|
|
||||||
|
// Callers get no handle to these until the Control is returned, release them on any error.
|
||||||
|
defer func() {
|
||||||
|
if reterr != nil {
|
||||||
|
for _, u := range udpConns {
|
||||||
|
if u != nil {
|
||||||
|
_ = u.Close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
if !configTest {
|
if !configTest {
|
||||||
rawListenHost := c.GetString("listen.host", "0.0.0.0")
|
rawListenHost := c.GetString("listen.host", "0.0.0.0")
|
||||||
var listenHost netip.Addr
|
var listenHost netip.Addr
|
||||||
|
|||||||
@@ -40,6 +40,7 @@ func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip
|
|||||||
|
|
||||||
err := t.reload(c, true)
|
err := t.reload(c, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
_ = file.Close()
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -659,7 +659,6 @@ func addRoute(prefix netip.Prefix, gateway netroute.Addr) error {
|
|||||||
return fmt.Errorf("failed to create route.RouteMessage for change: %w", err)
|
return fmt.Errorf("failed to create route.RouteMessage for change: %w", err)
|
||||||
}
|
}
|
||||||
_, err = unix.Write(sock, data[:])
|
_, err = unix.Write(sock, data[:])
|
||||||
fmt.Println("DOING CHANGE")
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return fmt.Errorf("failed to write route.RouteMessage to socket: %w", err)
|
return fmt.Errorf("failed to write route.RouteMessage to socket: %w", err)
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ import (
|
|||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
@@ -33,6 +34,12 @@ func newTun(_ *config.C, _ *slog.Logger, _ []netip.Prefix, _ bool) (*tun, error)
|
|||||||
}
|
}
|
||||||
|
|
||||||
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
|
if err := unix.SetNonblock(deviceFd, true); err != nil {
|
||||||
|
// We own the fd from the moment it is handed to us, same as the reload error path below
|
||||||
|
_ = unix.Close(deviceFd)
|
||||||
|
return nil, fmt.Errorf("failed to set the tun fd to non-blocking mode: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
file := os.NewFile(uintptr(deviceFd), "/dev/tun")
|
file := os.NewFile(uintptr(deviceFd), "/dev/tun")
|
||||||
t := &tun{
|
t := &tun{
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
@@ -42,6 +49,7 @@ func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip
|
|||||||
|
|
||||||
err := t.reload(c, true)
|
err := t.reload(c, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
_ = file.Close()
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+4
-35
@@ -5,7 +5,6 @@ package overlay
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"errors"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
@@ -251,17 +250,6 @@ func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip
|
|||||||
}
|
}
|
||||||
|
|
||||||
func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) {
|
func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) {
|
||||||
// Validate the device name up front so a bad tun.dev fails fast, before we
|
|
||||||
// open /dev/net/tun or leak a file descriptor. A single %d in the name is
|
|
||||||
// substituted by the kernel during TUNSETIFF (dev_alloc_name) with the
|
|
||||||
// lowest number that yields an unused device name. Resolving the template
|
|
||||||
// in the kernel keeps the pick-a-name/create-the-device pair atomic, so
|
|
||||||
// concurrent callers can never race each other to the same name.
|
|
||||||
tunName := c.GetString("tun.dev", "nebula%d")
|
|
||||||
if err := validateTunName(tunName); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// If /dev/net/tun doesn't exist, try to create it (will happen in docker)
|
// If /dev/net/tun doesn't exist, try to create it (will happen in docker)
|
||||||
@@ -289,11 +277,12 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
|
|||||||
if multiqueue {
|
if multiqueue {
|
||||||
req.Flags |= unix.IFF_MULTI_QUEUE
|
req.Flags |= unix.IFF_MULTI_QUEUE
|
||||||
}
|
}
|
||||||
copy(req.Name[:], tunName)
|
nameStr := c.GetString("tun.dev", "")
|
||||||
|
copy(req.Name[:], nameStr)
|
||||||
if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
||||||
_ = unix.Close(fd)
|
_ = unix.Close(fd)
|
||||||
return nil, &NameError{
|
return nil, &NameError{
|
||||||
Name: tunName,
|
Name: nameStr,
|
||||||
Underlying: err,
|
Underlying: err,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -309,27 +298,6 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
|
|||||||
return t, nil
|
return t, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func validateTunName(tunName string) error {
|
|
||||||
if !strings.Contains(tunName, "%d") {
|
|
||||||
if len(tunName) >= unix.IFNAMSIZ {
|
|
||||||
return fmt.Errorf("tun.dev %q is not shorter than the maximum device name length of %d", tunName, unix.IFNAMSIZ)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
if strings.Count(tunName, "%d") > 1 {
|
|
||||||
return fmt.Errorf("tun.dev template %q may only contain a single %%d", tunName)
|
|
||||||
}
|
|
||||||
if tunName == "%d" {
|
|
||||||
return errors.New("please don't name your tun device '%d'")
|
|
||||||
}
|
|
||||||
// The kernel substitutes the %d itself and requires the template, like a
|
|
||||||
// literal name, to be NUL-terminated within IFNAMSIZ bytes.
|
|
||||||
if len(tunName) >= unix.IFNAMSIZ {
|
|
||||||
return fmt.Errorf("tun.dev template %q is not shorter than the maximum device name length of %d", tunName, unix.IFNAMSIZ)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// newTunGeneric does all the stuff common to different tun initialization paths. It will close your files on error.
|
// newTunGeneric does all the stuff common to different tun initialization paths. It will close your files on error.
|
||||||
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
tfd, err := newTunFd(fd)
|
tfd, err := newTunFd(fd)
|
||||||
@@ -800,6 +768,7 @@ func (t *tun) isGatewayInVpnNetworks(gwAddr netip.Addr) bool {
|
|||||||
|
|
||||||
func (t *tun) getGatewaysFromRoute(r *netlink.Route) routing.Gateways {
|
func (t *tun) getGatewaysFromRoute(r *netlink.Route) routing.Gateways {
|
||||||
var gateways routing.Gateways
|
var gateways routing.Gateways
|
||||||
|
|
||||||
link, err := netlink.LinkByName(t.Device)
|
link, err := netlink.LinkByName(t.Device)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.l.Error("Ignoring route update: failed to get link by name", "deviceName", t.Device)
|
t.l.Error("Ignoring route update: failed to get link by name", "deviceName", t.Device)
|
||||||
|
|||||||
@@ -3,12 +3,7 @@
|
|||||||
|
|
||||||
package overlay
|
package overlay
|
||||||
|
|
||||||
import (
|
import "testing"
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
|
||||||
)
|
|
||||||
|
|
||||||
var runAdvMSSTests = []struct {
|
var runAdvMSSTests = []struct {
|
||||||
name string
|
name string
|
||||||
@@ -37,39 +32,3 @@ func TestTunAdvMSS(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestValidateTunName(t *testing.T) {
|
|
||||||
// A device name must be shorter than IFNAMSIZ (i.e. IFNAMSIZ-1 chars max).
|
|
||||||
maxLenName := strings.Repeat("a", unix.IFNAMSIZ-1)
|
|
||||||
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
tmpl string
|
|
||||||
wantErr bool
|
|
||||||
}{
|
|
||||||
{"short literal name is fine", "nebula1", false},
|
|
||||||
{"literal name at the max length is fine", maxLenName, false},
|
|
||||||
{"literal name at IFNAMSIZ is rejected", strings.Repeat("a", unix.IFNAMSIZ), true},
|
|
||||||
{"trailing template is fine", "nebula%d", false},
|
|
||||||
{"mid-string template is fine", "neb%dprod", false},
|
|
||||||
{"leading template is fine", "%dnebula", false},
|
|
||||||
{"template at the max length is fine", strings.Repeat("a", unix.IFNAMSIZ-3) + "%d", false},
|
|
||||||
{"template at IFNAMSIZ is rejected", strings.Repeat("a", unix.IFNAMSIZ-2) + "%d", true},
|
|
||||||
{"bare %d is rejected", "%d", true},
|
|
||||||
{"multiple %d is rejected", "neb%d%dprod", true},
|
|
||||||
{"over-long template is rejected", strings.Repeat("a", unix.IFNAMSIZ-1) + "%d", true},
|
|
||||||
{"over-long mid-string template is rejected", "neb%d" + strings.Repeat("a", unix.IFNAMSIZ-3), true},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
err := validateTunName(tt.tmpl)
|
|
||||||
if tt.wantErr && err == nil {
|
|
||||||
t.Fatalf("expected an error for %q, got none", tt.tmpl)
|
|
||||||
}
|
|
||||||
if !tt.wantErr && err != nil {
|
|
||||||
t.Fatalf("unexpected error for %q: %v", tt.tmpl, err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
+9
-1
@@ -107,7 +107,10 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak
|
|||||||
if relayHostInfo.GetRemote().IsValid() {
|
if relayHostInfo.GetRemote().IsValid() {
|
||||||
idx, err := AddRelay(rm.l, relayHostInfo, rm.hostmap, vpnIp, nil, TerminalType, Requested)
|
idx, err := AddRelay(rm.l, relayHostInfo, rm.hostmap, vpnIp, nil, TerminalType, Requested)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
// No local relay state was installed, so a CreateRelayRequest would hand the
|
||||||
|
// peer an index we could never resolve. Skip it.
|
||||||
hl.Info("Failed to add relay to hostmap", "relay", relay.String(), "error", err)
|
hl.Info("Failed to add relay to hostmap", "relay", relay.String(), "error", err)
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
m := NebulaControl{
|
m := NebulaControl{
|
||||||
@@ -237,7 +240,12 @@ func AddRelay(l *slog.Logger, relayHostInfo *HostInfo, hm *HostMap, vpnIp netip.
|
|||||||
// Avoid standing up a relay that can't be used since only the primary hostinfo
|
// Avoid standing up a relay that can't be used since only the primary hostinfo
|
||||||
// will be pointed to by the relay logic
|
// will be pointed to by the relay logic
|
||||||
//TODO: if there was an existing primary and it had relay state, should we merge?
|
//TODO: if there was an existing primary and it had relay state, should we merge?
|
||||||
hm.unlockedMakePrimary(relayHostInfo)
|
if !hm.unlockedMakePrimary(relayHostInfo) {
|
||||||
|
// The tunnel was torn down after the caller grabbed relayHostInfo. A relay standing
|
||||||
|
// on an unlinked hostinfo would never carry traffic, and its Relays entry could
|
||||||
|
// never be reclaimed since the delete-time cleanup has already run.
|
||||||
|
return 0, errors.New("relay hostinfo is no longer in the hostmap")
|
||||||
|
}
|
||||||
|
|
||||||
hm.Relays[index] = relayHostInfo
|
hm.Relays[index] = relayHostInfo
|
||||||
newRelay := Relay{
|
newRelay := Relay{
|
||||||
|
|||||||
+16
-8
@@ -43,12 +43,25 @@ type Service struct {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func New(control *nebula.Control) (*Service, error) {
|
func New(control *nebula.Control) (_ *Service, reterr error) {
|
||||||
wait, err := control.Start()
|
// Check this before Start so a failure doesn't leave a running nebula
|
||||||
|
device, ok := control.Device().(*overlay.UserDevice)
|
||||||
|
if !ok {
|
||||||
|
return nil, errors.New("must be using user device")
|
||||||
|
}
|
||||||
|
|
||||||
|
err := control.Start()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Anything that fails after a successful Start must tear nebula back down
|
||||||
|
defer func() {
|
||||||
|
if reterr != nil {
|
||||||
|
control.Stop()
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
ctx := control.Context()
|
ctx := control.Context()
|
||||||
eg, ctx := errgroup.WithContext(ctx)
|
eg, ctx := errgroup.WithContext(ctx)
|
||||||
s := Service{
|
s := Service{
|
||||||
@@ -57,11 +70,6 @@ func New(control *nebula.Control) (*Service, error) {
|
|||||||
}
|
}
|
||||||
s.mu.listeners = map[uint16]*tcpListener{}
|
s.mu.listeners = map[uint16]*tcpListener{}
|
||||||
|
|
||||||
device, ok := control.Device().(*overlay.UserDevice)
|
|
||||||
if !ok {
|
|
||||||
return nil, errors.New("must be using user device")
|
|
||||||
}
|
|
||||||
|
|
||||||
s.ipstack = stack.New(stack.Options{
|
s.ipstack = stack.New(stack.Options{
|
||||||
NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol},
|
NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol},
|
||||||
TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol, udp.NewProtocol, icmp.NewProtocol4, icmp.NewProtocol6},
|
TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol, udp.NewProtocol, icmp.NewProtocol4, icmp.NewProtocol6},
|
||||||
@@ -147,7 +155,7 @@ func New(control *nebula.Control) (*Service, error) {
|
|||||||
// Add the nebula wait function to the group so a fatal reader error
|
// Add the nebula wait function to the group so a fatal reader error
|
||||||
// propagates out through errgroup.Wait().
|
// propagates out through errgroup.Wait().
|
||||||
eg.Go(func() error {
|
eg.Go(func() error {
|
||||||
return wait()
|
return control.Wait()
|
||||||
})
|
})
|
||||||
|
|
||||||
return &s, nil
|
return &s, nil
|
||||||
|
|||||||
Reference in New Issue
Block a user