Compare commits

..

2 Commits

Author SHA1 Message Date
Nate Brown 35ac12b0c0 Group the HostInfo fields the packet paths touch 2026-07-23 21:31:51 -05:00
Nate Brown c07f28cd04 Fold the rebind counter and traffic flags into one atomic word 2026-07-23 20:36:51 -05:00
14 changed files with 257 additions and 966 deletions
+13 -8
View File
@@ -105,11 +105,17 @@ func (cm *connectionManager) getInactivityTimeout() time.Duration {
}
func (cm *connectionManager) In(h *HostInfo) {
h.in.Store(true)
h.markIn()
}
func (cm *connectionManager) Out(h *HostInfo) {
h.out.Store(true)
// OutRelay records relayed traffic, leaving the rebind epoch for the direct path to this host to consume
func (cm *connectionManager) OutRelay(h *HostInfo) {
h.markOutOnly()
}
// Out records outbound traffic and reports whether we rebound since this tunnel last sent
func (cm *connectionManager) Out(h *HostInfo) bool {
return h.markOut(cm.intf.rebindEpoch.Load())
}
func (cm *connectionManager) RelayUsed(localIndex uint32) {
@@ -128,8 +134,7 @@ func (cm *connectionManager) RelayUsed(localIndex uint32) {
// getAndResetTrafficCheck returns if there was any inbound or outbound traffic within the last tick and
// resets the state for this local index
func (cm *connectionManager) getAndResetTrafficCheck(h *HostInfo, now time.Time) (bool, bool) {
in := h.in.Swap(false)
out := h.out.Swap(false)
in, out := h.takeTraffic()
if in || out {
h.lastUsed = now
}
@@ -340,7 +345,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
"tunnelCheck", m{"state": "alive", "method": "passive"},
)
}
hostinfo.pendingDeletion.Store(false)
hostinfo.setPendingDeletion(false)
if mainHostInfo {
decision = tryRehandshake
@@ -363,7 +368,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
return decision, hostinfo, primary
}
if hostinfo.pendingDeletion.Load() {
if hostinfo.isPendingDeletion() {
// We have already sent a test packet and nothing was returned, this hostinfo is dead
hostinfo.logger(cm.l).Info("Tunnel status",
"tunnelCheck", m{"state": "dead", "method": "active"},
@@ -414,7 +419,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
}
}
hostinfo.pendingDeletion.Store(true)
hostinfo.setPendingDeletion(true)
cm.trafficTimer.Add(hostinfo.localIndexId, cm.pendingDeletionInterval)
return decision, hostinfo, nil
}
+36 -36
View File
@@ -86,25 +86,25 @@ func Test_NewConnectionManagerTest(t *testing.T) {
// We saw traffic out to vpnIp
nc.Out(hostinfo)
nc.In(hostinfo)
assert.False(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.isPendingDeletion())
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
assert.True(t, hostinfo.out.Load())
assert.True(t, hostinfo.in.Load())
assert.True(t, hostinfo.sentSinceCheck())
assert.True(t, (hostinfo.state.Load()&stateIn != 0))
// Do a traffic check tick, should not be pending deletion but should not have any in/out packets recorded
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
assert.False(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
assert.False(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
// Do another traffic check tick, this host should be pending deletion now
nc.Out(hostinfo)
assert.True(t, hostinfo.out.Load())
assert.True(t, hostinfo.sentSinceCheck())
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
assert.True(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
assert.True(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
@@ -168,33 +168,33 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
// We saw traffic out to vpnIp
nc.Out(hostinfo)
nc.In(hostinfo)
assert.True(t, hostinfo.in.Load())
assert.True(t, hostinfo.out.Load())
assert.False(t, hostinfo.pendingDeletion.Load())
assert.True(t, (hostinfo.state.Load()&stateIn != 0))
assert.True(t, hostinfo.sentSinceCheck())
assert.False(t, hostinfo.isPendingDeletion())
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
// Do a traffic check tick, should not be pending deletion but should not have any in/out packets recorded
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
assert.False(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
assert.False(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
// Do another traffic check tick, this host should be pending deletion now
nc.Out(hostinfo)
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
assert.True(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
assert.True(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
// We saw traffic, should no longer be pending deletion
nc.In(hostinfo)
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
assert.False(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
assert.False(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
}
@@ -253,31 +253,31 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
// Do a traffic check tick, in and out should be cleared but should not be pending deletion
nc.Out(hostinfo)
nc.In(hostinfo)
assert.True(t, hostinfo.out.Load())
assert.True(t, hostinfo.in.Load())
assert.True(t, hostinfo.sentSinceCheck())
assert.True(t, (hostinfo.state.Load()&stateIn != 0))
now := time.Now()
decision, _, _ := nc.makeTrafficDecision(hostinfo.localIndexId, now)
assert.Equal(t, tryRehandshake, decision)
assert.Equal(t, now, hostinfo.lastUsed)
assert.False(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
assert.False(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Second*5))
assert.Equal(t, doNothing, decision)
assert.Equal(t, now, hostinfo.lastUsed)
assert.False(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
assert.False(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
// Do another traffic check tick, should still not be pending deletion
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Second*10))
assert.Equal(t, doNothing, decision)
assert.Equal(t, now, hostinfo.lastUsed)
assert.False(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
assert.False(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
@@ -285,9 +285,9 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Minute*10))
assert.Equal(t, closeTunnel, decision)
assert.Equal(t, now, hostinfo.lastUsed)
assert.False(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
assert.False(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
}
+1 -5
View File
@@ -52,7 +52,6 @@ type Control struct {
sshStart func()
statsStart func()
dnsStart func()
infoAPIStart func()
lighthouseStart func()
networkChangeStart func(rebind func())
connectionManagerStart func(context.Context)
@@ -109,9 +108,6 @@ func (c *Control) Start() error {
if c.networkChangeStart != nil {
go c.networkChangeStart(c.RebindUDPServer)
}
if c.infoAPIStart != nil {
go c.infoAPIStart()
}
if c.connectionManagerStart != nil {
go c.connectionManagerStart(c.ctx)
}
@@ -216,7 +212,7 @@ func (c *Control) RebindUDPServer() {
c.f.lightHouse.SendUpdate()
// Let the main interface know that we rebound so that underlying tunnels know to trigger punches from their remotes
c.f.rebindCount++
c.f.rebindEpoch.Add(1)
}
// ListHostmapHosts returns details about the actual or pending (handshaking) hostmap by vpn ip
+10
View File
@@ -123,6 +123,16 @@ func (c *Control) SetLocalAddrsFn(fn func(*LocalAllowList) []netip.Addr) {
c.f.lightHouse.localAddrsFn = fn
}
// GetRebindEpochFor returns the rebind epoch a tunnel last sent under, so a test can tell whether a send
// consumed the epoch edge without having to infer it from lighthouse traffic.
func (c *Control) GetRebindEpochFor(vpnAddr netip.Addr) (uint32, bool) {
h := c.f.hostMap.QueryVpnAddr(vpnAddr)
if h == nil {
return 0, false
}
return h.state.Load() >> stateEpochShift, true
}
func (c *Control) KillPendingTunnel(vpnIp netip.Addr) bool {
hostinfo := c.f.handshakeManager.QueryVpnAddr(vpnIp)
if hostinfo == nil {
+22 -3
View File
@@ -258,12 +258,31 @@ func (d *dnsServer) QueryCert(data string) string {
return ""
}
crt := findCertificateForVpnAddr(d.certState(), d.hostMap, ip)
if crt == nil {
// The hostmap only ever contains peers we have handshaked with, so it never carries an entry for ourselves.
// Answer self lookups straight from the local cert state.
if cs := d.certState(); cs != nil && cs.myVpnAddrsTable != nil && cs.myVpnAddrsTable.Contains(ip) {
c := cs.GetDefaultCertificate()
if c == nil {
return ""
}
b, err := c.MarshalJSON()
if err != nil {
return ""
}
return string(b)
}
hostinfo := d.hostMap.QueryVpnAddr(ip)
if hostinfo == nil {
return ""
}
b, err := crt.MarshalJSON()
q := hostinfo.GetCert()
if q == nil {
return ""
}
b, err := q.Certificate.MarshalJSON()
if err != nil {
return ""
}
+54
View File
@@ -223,3 +223,57 @@ func TestRebindAdvertisesNewAddressAfterMove(t *testing.T) {
lhControl.Stop()
myControl.Stop()
}
// A relayed send records traffic but must not consume the rebind epoch. If it does, the next direct send to the
// relay host sees the epoch already current and never requeries, so the far side is never told to punch at our
// new address. This pins the SendVia call site, which the unit tests cannot reach.
func TestRebindRequeriesAfterRelayedSend(t *testing.T) {
t.Parallel()
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
// No lighthouse on purpose: it would hand out a direct address for them and nothing would relay.
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.2/24", m{"relay": m{"use_relays": true}})
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
r := router.NewR(t, myControl, relayControl, theirControl)
defer r.RenderFlow()
myControl.Start()
relayControl.Start()
theirControl.Start()
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("establish")))
r.RouteForAllUntilTxTun(theirControl)
r.RouteFor(time.Millisecond * 500)
hi := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false)
require.NotNil(t, hi, "expected a tunnel to them")
require.NotEmpty(t, hi.CurrentRelaysToMe, "them must be reachable only via the relay for this test to mean anything")
// sendNoMetrics only reaches SendVia when there is no direct remote, so pin that too. Without this the test
// keeps passing while quietly sending direct and never exercising the relay path.
require.False(t, hi.CurrentRemote.IsValid(), "them must have no direct remote, otherwise SendVia is never called")
before, ok := myControl.GetRebindEpochFor(relayVpnIpNet[0].Addr())
require.True(t, ok, "expected a tunnel to the relay")
myControl.RebindUDPServer()
// Traffic to them goes through SendVia on the relay tunnel. That must record traffic without consuming the
// relay tunnel's own epoch edge, which belongs to the direct path.
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("relayed")))
r.RouteForAllUntilTxTun(theirControl)
after, ok := myControl.GetRebindEpochFor(relayVpnIpNet[0].Addr())
require.True(t, ok)
assert.Equal(t, before, after,
"a relayed send consumed the relay tunnel's rebind epoch, so the next direct send will not requery")
myControl.Stop()
relayControl.Stop()
theirControl.Stop()
}
-29
View File
@@ -231,35 +231,6 @@ punchy:
# Overriding this to "" is the same as "/" and will allow overwriting any path on the host.
#sandbox_dir: /var/tmp/nebula-debug
# EXPERIMENTAL: this feature may change or disappear in the future.
# info_api exposes a small local HTTP+JSON API that lets other programs on
# this machine resolve a vpn address to its certificate identity (name, vpn
# addresses, groups, fingerprint, validity), e.g. for making authorization
# decisions about an inbound connection:
# GET /v1/host?addr=<vpn addr> - identity of the host owning the address: a
# peer with an active tunnel, or this node itself. `addr` may include a
# port (`192.168.100.7:54321`), which is ignored, so a connection's remote
# address can be passed through as is. Returns 404 when the address is
# unknown or has no active tunnel.
# GET /v1/self - this node's own identity.
# Identity answers can be trusted because nebula drops inbound packets whose
# source vpn address is not contained in the sender's certificate, so the
# source address of a connection arriving over the nebula interface is
# guaranteed to map to the certificate reported here.
# There is no authentication in this API; restrict access with unix socket
# file permissions.
# This whole section is reloadable.
#info_api:
# Toggles the feature
#enabled: false
# listen accepts a unix socket path as a unix:// URL with an absolute path:
#listen: unix:///var/run/nebula-info-api.sock
# File mode for the unix socket, as an octal string.
# The socket is created by nebula's user; to grant a group of local services
# access, place the socket in a directory with appropriate permissions
# (e.g. a systemd RuntimeDirectory) and relax this to "0660".
#socket_mode: "0600"
# EXPERIMENTAL: relay support for networks that can't establish direct connections.
relay:
# Relays are a list of Nebula IP's that peers can use to relay packets to me.
+73 -11
View File
@@ -238,18 +238,26 @@ const (
)
type HostInfo struct {
// The first cache line is everything the packet paths touch.
remote atomic.Pointer[netip.AddrPort]
remotes *RemoteList
promoteCounter atomic.Uint32
ConnectionState *ConnectionState
// Traffic bits, pendingDeletion, and the rebind epoch we last sent under
state atomic.Uint32
promoteCounter atomic.Uint32
remoteIndexId uint32
localIndexId uint32
remotes *RemoteList
// vpnAddrs is a list of vpn addresses assigned to this host that are within our own vpn networks
// The host may have other vpn addresses that are outside our
// vpn networks but were removed because they are not usable
vpnAddrs []netip.Addr
// Everything below is off the packet path: handshakes, relays, roaming and the connection manager.
// networks is a combination of specific vpn addresses (not prefixes!) and full unsafe networks assigned to this host.
networks *bart.Table[NetworkType]
relayState RelayState
@@ -262,11 +270,6 @@ type HostInfo struct {
// This is used to limit lighthouse re-queries in chatty clients
nextLHQuery atomic.Int64
// lastRebindCount is the other side of Interface.rebindCount, if these values don't match then we need to ask LH
// for a punch from the remote end of this tunnel. The goal being to prime their conntrack for our traffic just like
// with a handshake
lastRebindCount int8
// lastHandshakeTime records the time the remote side told us about at the stage when the handshake was completed locally
// Stage 1 packet will contain it if I am a responder, stage 2 packet if I am an initiator
// This is used to avoid an attack where a handshake packet is replayed after some time
@@ -275,9 +278,6 @@ type HostInfo struct {
lastRoam time.Time
lastRoamRemote netip.AddrPort
//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
// lastUsed tracks the last time ConnectionManager checked the tunnel and it was in use.
// This value will be behind against actual tunnel utilization in the hot path.
// This should only be used by the ConnectionManagers ticker routine.
@@ -658,7 +658,7 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
hm.Indexes[hostinfo.localIndexId] = hostinfo
hm.RemoteIndexes[hostinfo.remoteIndexId] = hostinfo
hostinfo.out.Store(true)
hostinfo.markOut(f.rebindEpoch.Load())
if f.connectionManager != nil { // f.connectionManager is only nil in some unit tests
f.connectionManager.trafficTimer.Add(hostinfo.localIndexId, f.connectionManager.checkInterval)
}
@@ -759,6 +759,68 @@ func (i *HostInfo) TryPromoteBest(preferredRanges []netip.Prefix, ifce *Interfac
}
}
// Bits within HostInfo.state, everything above stateEpochShift is the epoch
const (
stateIn uint32 = 1 << iota
stateOut
statePendingDeletion
stateFlags = stateIn | stateOut | statePendingDeletion
stateEpochShift = 3
)
// markIn records inbound traffic
func (i *HostInfo) markIn() {
if i.state.Load()&stateIn == 0 {
i.state.Or(stateIn)
}
}
// markOut records a send and reports whether the epoch moved, meaning we want a punch from the far side
func (i *HostInfo) markOut(epoch uint32) bool {
e := epoch << stateEpochShift
for {
old := i.state.Load()
if old&stateOut != 0 && old&^stateFlags == e {
return false
}
if i.state.CompareAndSwap(old, old&stateFlags|stateOut|e) {
return old&^stateFlags != e
}
}
}
// markOutOnly records a send without consuming the rebind epoch, for paths that cannot act on a requery
func (i *HostInfo) markOutOnly() {
if i.state.Load()&stateOut == 0 {
i.state.Or(stateOut)
}
}
// sentSinceCheck reports whether anything has been sent since the connection manager last looked
func (i *HostInfo) sentSinceCheck() bool {
return i.state.Load()&stateOut != 0
}
// takeTraffic clears both traffic bits, leaving the epoch alone, and reports what they were
func (i *HostInfo) takeTraffic() (in bool, out bool) {
old := i.state.And(^(stateIn | stateOut))
return old&stateIn != 0, old&stateOut != 0
}
func (i *HostInfo) setPendingDeletion(v bool) {
if v {
i.state.Or(statePendingDeletion)
} else {
i.state.And(^statePendingDeletion)
}
}
func (i *HostInfo) isPendingDeletion() bool {
return i.state.Load()&statePendingDeletion != 0
}
func (i *HostInfo) GetCert() *cert.CachedCertificate {
if i.ConnectionState != nil {
return i.ConnectionState.peerCert
+40
View File
@@ -401,3 +401,43 @@ func TestHostMap_RelayState(t *testing.T) {
assert.Equal(t, []netip.Addr{}, h1.relayState.relays)
}
func TestHostInfo_markOut(t *testing.T) {
h := &HostInfo{}
h.markOut(5) // stamped when the tunnel was added
// A tunnel already on the current epoch has nothing to report, which is what keeps a fresh tunnel from
// requerying on its first packet
assert.False(t, h.markOut(5), "an unchanged epoch should not report a move")
assert.True(t, h.sentSinceCheck(), "the send is still recorded as traffic")
// A rebind is observed exactly once, so we requery once per rebind
assert.True(t, h.markOut(6), "a bumped epoch should report a move")
assert.False(t, h.markOut(6), "the epoch move should only be reported once")
// Traffic and pendingDeletion live in the same word and must survive an epoch change
h.setPendingDeletion(true)
h.markIn()
assert.True(t, h.markOut(7))
assert.True(t, h.isPendingDeletion(), "pendingDeletion must survive an epoch change")
in, out := h.takeTraffic()
assert.True(t, in, "inbound traffic must survive an epoch change")
assert.True(t, out)
// Clearing the traffic bits leaves the epoch alone, otherwise an idle tunnel would requery forever
assert.False(t, h.markOut(7), "takeTraffic must not disturb the epoch")
}
// A relayed send records traffic but must leave the rebind epoch for the direct path to consume, otherwise
// relaying to a host swallows the requery that gets the far side punching at our new address.
func TestHostInfo_markOutOnly(t *testing.T) {
h := &HostInfo{}
h.markOut(5)
h.markOutOnly()
assert.True(t, h.sentSinceCheck(), "a relayed send is still outbound traffic")
assert.False(t, h.markOut(5), "a relayed send must not disturb the epoch")
assert.True(t, h.markOut(6), "a relayed send must not consume the epoch edge")
assert.False(t, h.markOut(6))
}
-409
View File
@@ -1,409 +0,0 @@
package nebula
import (
"context"
"encoding/json"
"errors"
"fmt"
"io/fs"
"log/slog"
"net"
"net/http"
"net/netip"
"os"
"path/filepath"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/config"
)
// infoAPIServer is a small http+json listener on a unix socket that lets other
// programs on this machine resolve a vpn address to its certificate identity (name, groups, networks)
// for making authorization decisions. Lifecycle works like statsServer: the constructor wires the
// reload callback, reload records config, Start runs the runtime, Stop tears it down
type infoAPIServer struct {
l *slog.Logger
ctx context.Context
hostMap *HostMap
pki *PKI
// enabled mirrors `info_api.enabled` so callers of Start don't need to know the gating rules
enabled atomic.Bool
runMu sync.Mutex
runCfg *infoAPIConfig
run *infoAPIRuntime // non-nil while a runtime is live
}
// infoAPIRuntime is the live state owned by a single Start invocation. Stop and Start's exit path
// use pointer equality to tell "my runtime" apart from one that replaced it after a reload
type infoAPIRuntime struct {
server *http.Server
listener net.Listener
}
// infoAPIConfig is a snapshot of the info_api config section, comparable with == so reload can
// detect "no change" cheaply
type infoAPIConfig struct {
enabled bool
listen string // raw config value, for error messages
addr string // unix socket path
// file mode applied to the unix socket after bind
socketMode fs.FileMode
}
// newInfoAPIServerFromConfig builds a infoAPIServer and applies the initial config. The reload
// callback is registered first so a SIGHUP can later enable, fix, or disable the listener even if
// the initial config was bad. Nothing binds until Start, so config tests are side effect free.
// A bad config is logged rather than returned: it must not stop nebula from starting, the feature
// just stays disabled until a reload provides a valid config
func newInfoAPIServerFromConfig(ctx context.Context, l *slog.Logger, pki *PKI, hostMap *HostMap, c *config.C) *infoAPIServer {
h := &infoAPIServer{
l: l,
ctx: ctx,
hostMap: hostMap,
pki: pki,
}
c.RegisterReloadCallback(func(c *config.C) {
if err := h.reload(c, false); err != nil {
h.l.Warn("Failed to reload info API from config", "error", err)
}
})
if err := h.reload(c, true); err != nil {
h.l.Warn("Failed to apply info API config; it will stay disabled until the config is fixed and reloaded", "error", err)
}
return h
}
// reload records the latest config. The initial call only records it, Control.Start launches the
// first runtime via infoAPIStart. Later calls reconcile the running listener with the new config:
// enable, disable, or restart when the listen config changed
func (h *infoAPIServer) reload(c *config.C, initial bool) error {
newCfg, err := loadInfoAPIConfig(c)
if err != nil {
return err
}
h.runMu.Lock()
sameCfg := h.runCfg != nil && *h.runCfg == newCfg
h.runCfg = &newCfg
running := h.run != nil
h.runMu.Unlock()
h.enabled.Store(newCfg.enabled)
if initial || sameCfg {
return nil
}
if running {
h.Stop()
}
if newCfg.enabled {
go h.Start()
}
return nil
}
// Start binds the listener from the latest config and serves until Stop is called or ctx fires.
// Safe to call when disabled or already running (both no-op)
func (h *infoAPIServer) Start() {
if !h.enabled.Load() {
return
}
h.runMu.Lock()
if h.ctx.Err() != nil || h.run != nil || h.runCfg == nil {
h.runMu.Unlock()
return
}
cfg := *h.runCfg
ln, err := h.listen(cfg)
if err != nil {
// drop the cached config so a SIGHUP with the same config retries the bind
h.runCfg = nil
h.runMu.Unlock()
h.l.Error("Failed to start info API listener", "listen", cfg.listen, "error", err)
return
}
mux := http.NewServeMux()
mux.HandleFunc("GET /v1/host", h.handleHost)
mux.HandleFunc("GET /v1/self", h.handleSelf)
srv := &http.Server{Handler: mux, ReadHeaderTimeout: 5 * time.Second}
rt := &infoAPIRuntime{server: srv, listener: ln}
h.run = rt
h.runMu.Unlock()
h.l.Info("Starting info API listener", "addr", ln.Addr())
cleanExit := h.serve(srv, ln)
// A Stop that raced our bind shut the server down before Serve could adopt the listener;
// closing it again is harmless and guarantees a unix socket file gets unlinked
_ = ln.Close()
// Clear our runtime only if nothing has replaced it. Stop races through here too but leaves
// h.run == nil, so the pointer check skips
h.runMu.Lock()
if h.run == rt {
h.run = nil
// an error exit leaves runCfg cached as if it were applied, drop it so a SIGHUP with the
// same config re-triggers Start once the user fixes the underlying problem
if !cleanExit {
h.runCfg = nil
}
}
h.runMu.Unlock()
}
// serve runs srv.Serve and ensures ctx cancellation unblocks it. Returns true if the listener
// exited cleanly (Stop, ctx cancellation), false on an unexpected error
func (h *infoAPIServer) serve(srv *http.Server, ln net.Listener) bool {
// ctx cancellation triggers a server shutdown which in turn unblocks Serve, closing `done` on
// exit keeps the watcher from outliving this call
done := make(chan struct{})
go func() {
select {
case <-h.ctx.Done():
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := srv.Shutdown(shutdownCtx); err != nil {
h.l.Warn("Failed to shut down info API listener", "error", err)
}
case <-done:
}
}()
defer close(done)
err := srv.Serve(ln)
if err == nil || errors.Is(err, http.ErrServerClosed) {
return true
}
h.l.Error("Info API listener exited", "error", err)
return false
}
// Stop tears down the active runtime, if any. Idempotent
func (h *infoAPIServer) Stop() {
h.runMu.Lock()
rt := h.run
h.run = nil
h.runMu.Unlock()
if rt == nil {
return
}
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := rt.server.Shutdown(shutdownCtx); err != nil {
h.l.Warn("Failed to shut down info API listener", "error", err)
}
}
// listen binds the configured unix socket. It also clears a stale socket file left by an unclean
// exit and applies the configured file mode
func (h *infoAPIServer) listen(cfg infoAPIConfig) (net.Listener, error) {
if fi, err := os.Stat(cfg.addr); err == nil {
if fi.Mode()&os.ModeSocket == 0 {
return nil, fmt.Errorf("info_api.listen path %s exists and is not a socket, refusing to replace it", cfg.addr)
}
// a normal shutdown unlinks the socket, so a file here means a previous process exited
// uncleanly, remove it so the bind below can succeed
if err = os.Remove(cfg.addr); err != nil {
return nil, fmt.Errorf("failed to remove stale socket %s: %w", cfg.addr, err)
}
}
ln, err := net.Listen("unix", cfg.addr)
if err != nil {
return nil, err
}
// The socket is briefly live with umask-derived permissions before this chmod lands, tolerated
// because connections accepted in that window still only reach this read-only API
if err = os.Chmod(cfg.addr, cfg.socketMode); err != nil {
_ = ln.Close()
return nil, fmt.Errorf("failed to set mode on socket %s: %w", cfg.addr, err)
}
return ln, nil
}
func (h *infoAPIServer) certState() *CertState {
if h.pki == nil {
return nil
}
return h.pki.getCertState()
}
// handleHost serves GET /v1/host?addr=<vpn addr>, answering with the identity of the host that
// owns the address: a peer with an active tunnel, or this node itself. addr may include a port,
// which is ignored, so clients can pass a connection's remote address through without parsing it
func (h *infoAPIServer) handleHost(w http.ResponseWriter, r *http.Request) {
q := r.URL.Query().Get("addr")
if q == "" {
writeJSONError(w, http.StatusBadRequest, "missing addr parameter")
return
}
ip, err := parseQueryAddrParam(q)
if err != nil {
writeJSONError(w, http.StatusBadRequest, "invalid address")
return
}
crt := findCertificateForVpnAddr(h.certState(), h.hostMap, ip)
if crt == nil {
writeJSONError(w, http.StatusNotFound, "no active tunnel for address")
return
}
h.writeHostIdentity(w, crt)
}
// handleSelf serves GET /v1/self, answering with this node's own identity
func (h *infoAPIServer) handleSelf(w http.ResponseWriter, r *http.Request) {
var crt cert.Certificate
if cs := h.certState(); cs != nil {
crt = cs.getCertificate(cs.initiatingVersion)
}
if crt == nil {
writeJSONError(w, http.StatusInternalServerError, "no certificate available")
return
}
h.writeHostIdentity(w, crt)
}
func (h *infoAPIServer) writeHostIdentity(w http.ResponseWriter, crt cert.Certificate) {
id, err := newHostIdentity(crt)
if err != nil {
writeJSONError(w, http.StatusInternalServerError, "failed to fingerprint certificate")
return
}
w.Header().Set("Content-Type", "application/json")
if err = json.NewEncoder(w).Encode(id); err != nil {
h.l.Debug("Failed to write info API response", "error", err)
}
}
// findCertificateForVpnAddr answers "who owns this vpn address": ourselves (from local cert state,
// the hostmap never carries an entry for this node) or a peer with an active tunnel. Returns nil
// when the address is unknown or the tunnel is mid-teardown
func findCertificateForVpnAddr(cs *CertState, hostMap *HostMap, ip netip.Addr) cert.Certificate {
if cs != nil && cs.myVpnAddrsTable != nil && cs.myVpnAddrsTable.Contains(ip) {
return cs.getCertificate(cs.initiatingVersion)
}
hostinfo := hostMap.QueryVpnAddr(ip)
if hostinfo == nil {
return nil
}
cc := hostinfo.GetCert()
if cc == nil {
return nil
}
return cc.Certificate
}
// hostIdentity is the json document served for both /v1/host and /v1/self, every field is derived
// from the authenticated certificate alone
type hostIdentity struct {
Name string `json:"name"`
VpnAddrs []netip.Addr `json:"vpnAddrs"`
Networks []netip.Prefix `json:"networks"`
UnsafeNetworks []netip.Prefix `json:"unsafeNetworks"`
Groups []string `json:"groups"`
Fingerprint string `json:"fingerprint"`
Issuer string `json:"issuer"`
NotBefore time.Time `json:"notBefore"`
NotAfter time.Time `json:"notAfter"`
CertVersion int `json:"certVersion"`
}
func newHostIdentity(crt cert.Certificate) (hostIdentity, error) {
fp, err := crt.Fingerprint()
if err != nil {
return hostIdentity{}, err
}
// slices are always allocated so they marshal as [] rather than null
networks := crt.Networks()
id := hostIdentity{
Name: crt.Name(),
VpnAddrs: make([]netip.Addr, 0, len(networks)),
Networks: append(make([]netip.Prefix, 0, len(networks)), networks...),
UnsafeNetworks: append(make([]netip.Prefix, 0, len(crt.UnsafeNetworks())), crt.UnsafeNetworks()...),
Groups: append(make([]string, 0, len(crt.Groups())), crt.Groups()...),
Fingerprint: fp,
Issuer: crt.Issuer(),
NotBefore: crt.NotBefore(),
NotAfter: crt.NotAfter(),
CertVersion: int(crt.Version()),
}
for _, n := range networks {
id.VpnAddrs = append(id.VpnAddrs, n.Addr())
}
return id, nil
}
func writeJSONError(w http.ResponseWriter, status int, msg string) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(map[string]string{"error": msg})
}
// parseQueryAddrParam parses the addr query parameter, accepting a bare address or an address with
// a port (`192.168.100.7:54321`, `[fd00::1]:443`) so callers can pass a connection's RemoteAddr
// straight through. The result is unmapped, 4in6 addresses (::ffff:a.b.c.d) become ipv4
func parseQueryAddrParam(s string) (netip.Addr, error) {
if ip, err := netip.ParseAddr(s); err == nil {
return ip.Unmap(), nil
}
ap, err := netip.ParseAddrPort(s)
if err != nil {
return netip.Addr{}, err
}
return ap.Addr().Unmap(), nil
}
func loadInfoAPIConfig(c *config.C) (infoAPIConfig, error) {
cfg := infoAPIConfig{
enabled: c.GetBool("info_api.enabled", false),
listen: c.GetString("info_api.listen", ""),
}
if !cfg.enabled {
return cfg, nil
}
if cfg.listen == "" {
return cfg, errors.New("info_api.listen can not be empty when info_api is enabled")
}
addr, err := parseInfoAPIListen(cfg.listen)
if err != nil {
return cfg, err
}
cfg.addr = addr
// read as a string so yaml can't reinterpret the octal literal
modeStr := c.GetString("info_api.socket_mode", "0600")
mode, err := strconv.ParseUint(modeStr, 8, 32)
if err != nil || fs.FileMode(mode)&^fs.ModePerm != 0 {
return cfg, fmt.Errorf("info_api.socket_mode was not a valid octal file mode: %s", modeStr)
}
cfg.socketMode = fs.FileMode(mode)
return cfg, nil
}
// parseInfoAPIListen extracts the unix socket path from the info_api.listen config value, which
// must be a `unix://` URL with an absolute path, e.g. `unix:///var/run/nebula.sock`
func parseInfoAPIListen(listen string) (addr string, err error) {
path, ok := strings.CutPrefix(listen, "unix://")
if !ok {
return "", fmt.Errorf("info_api.listen must be a unix:// socket path: %s", listen)
} else if !filepath.IsAbs(path) {
return "", fmt.Errorf("info_api.listen unix socket path must be absolute: %s", listen)
}
return path, nil
}
-448
View File
@@ -1,448 +0,0 @@
package nebula
import (
"context"
"encoding/json"
"fmt"
"io/fs"
"log/slog"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"net/url"
"os"
"path/filepath"
"runtime"
"testing"
"time"
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/cert_test"
"github.com/slackhq/nebula/config"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func Test_parseInfoAPIListen(t *testing.T) {
type testCase struct {
listen string
addr string
wantErr bool
}
tests := []testCase{
{listen: "", wantErr: true},
{listen: "unix://", wantErr: true},
{listen: "unix://relative/path.sock", wantErr: true},
{listen: "not an address", wantErr: true},
// tcp host:port addresses are no longer accepted
{listen: "127.0.0.1:8085", wantErr: true},
{listen: "[::1]:8085", wantErr: true},
{listen: "localhost:8085", wantErr: true},
}
// A unix socket path must be absolute for the OS that will bind it, and filepath.IsAbs is
// GOOS-specific. CI runs the suite separately on each OS, so assert the platform's own native
// absolute path is accepted while the other platform's is rejected.
posixPath := "unix:///var/run/nebula.sock"
winPath := `unix://C:\nebula\hq.sock`
if runtime.GOOS == "windows" {
tests = append(tests,
testCase{listen: winPath, addr: `C:\nebula\hq.sock`},
testCase{listen: posixPath, wantErr: true},
)
} else {
tests = append(tests,
testCase{listen: posixPath, addr: "/var/run/nebula.sock"},
testCase{listen: winPath, wantErr: true},
)
}
for _, tt := range tests {
addr, err := parseInfoAPIListen(tt.listen)
if tt.wantErr {
require.Error(t, err, "listen=%q", tt.listen)
continue
}
require.NoError(t, err, "listen=%q", tt.listen)
assert.Equal(t, tt.addr, addr, "listen=%q", tt.listen)
}
}
func Test_loadInfoAPIConfig(t *testing.T) {
c := config.NewC(nil)
// the listen path must be absolute for the OS running the test (CI is per-OS)
listen, wantAddr := "unix:///tmp/hq.sock", "/tmp/hq.sock"
if runtime.GOOS == "windows" {
listen, wantAddr = `unix://C:\tmp\hq.sock`, `C:\tmp\hq.sock`
}
// absent section means disabled, no error
cfg, err := loadInfoAPIConfig(c)
require.NoError(t, err)
assert.False(t, cfg.enabled)
// enabled without a listen address is an error
setInfoAPIConfig(c, true, "", "")
_, err = loadInfoAPIConfig(c)
require.Error(t, err)
// a unix socket gets the default mode
setInfoAPIConfig(c, true, listen, "")
cfg, err = loadInfoAPIConfig(c)
require.NoError(t, err)
assert.Equal(t, wantAddr, cfg.addr)
assert.Equal(t, fs.FileMode(0o600), cfg.socketMode)
setInfoAPIConfig(c, true, listen, "0660")
cfg, err = loadInfoAPIConfig(c)
require.NoError(t, err)
assert.Equal(t, fs.FileMode(0o660), cfg.socketMode)
setInfoAPIConfig(c, true, listen, "withers")
_, err = loadInfoAPIConfig(c)
require.Error(t, err)
// mode bits beyond the permission bits are rejected
setInfoAPIConfig(c, true, listen, "10600")
_, err = loadInfoAPIConfig(c)
require.Error(t, err)
// tcp host:port listen addresses are no longer supported
setInfoAPIConfig(c, true, "127.0.0.1:8085", "")
_, err = loadInfoAPIConfig(c)
require.Error(t, err)
}
func TestInfoAPIServer_badConfigIsNonFatal(t *testing.T) {
// an enabled-but-invalid config must not stop construction; nebula keeps starting and the
// feature simply stays disabled until a reload supplies a valid config
c := config.NewC(nil)
setInfoAPIConfig(c, true, "not-a-unix-socket", "")
h := newInfoAPIServerFromConfig(context.Background(), slog.New(slog.DiscardHandler), nil, newHostMap(slog.New(slog.DiscardHandler)), c)
require.NotNil(t, h)
assert.False(t, h.enabled.Load())
// no config was recorded, so Start has nothing to bind and is a no-op
h.runMu.Lock()
assert.Nil(t, h.runCfg)
h.runMu.Unlock()
h.Start()
h.runMu.Lock()
assert.Nil(t, h.run)
h.runMu.Unlock()
}
func setInfoAPIConfig(c *config.C, enabled bool, listen, socketMode string) {
settings := map[string]any{
"enabled": enabled,
"listen": listen,
}
if socketMode != "" {
settings["socket_mode"] = socketMode
}
c.Settings["info_api"] = settings
}
func newTestInfoAPIServer(t *testing.T) (*infoAPIServer, *config.C) {
t.Helper()
h := &infoAPIServer{
l: slog.New(slog.DiscardHandler),
ctx: context.Background(),
hostMap: newHostMap(slog.New(slog.DiscardHandler)),
}
h.hostMap.preferredRanges.Store(&[]netip.Prefix{})
return h, config.NewC(nil)
}
// addTestPeer creates a certificate for a peer owning each addr (as a /24 or /64) and inserts it
// into the hostmap as an established tunnel
func addTestPeer(t *testing.T, hm *HostMap, name string, addrs []netip.Addr, unsafeNetworks []netip.Prefix, groups []string) cert.Certificate {
t.Helper()
networks := make([]netip.Prefix, 0, len(addrs))
for _, a := range addrs {
bits := 24
if a.Is6() {
bits = 64
}
networks = append(networks, netip.PrefixFrom(a, bits))
}
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil)
crt, _, _, _ := cert_test.NewTestCert(cert.Version2, cert.Curve_CURVE25519, ca, caKey, name, time.Time{}, time.Time{}, networks, unsafeNetworks, groups)
fp, err := crt.Fingerprint()
require.NoError(t, err)
hm.unlockedAddHostInfo(&HostInfo{
ConnectionState: &ConnectionState{
peerCert: &cert.CachedCertificate{Certificate: crt, Fingerprint: fp},
},
vpnAddrs: addrs,
relayState: RelayState{
relayForByAddr: map[netip.Addr]*Relay{},
relayForByIdx: map[uint32]*Relay{},
},
}, &Interface{})
return crt
}
func getHost(t *testing.T, h *infoAPIServer, addrParam string) (int, map[string]any) {
t.Helper()
r := httptest.NewRequest(http.MethodGet, "/v1/host?addr="+url.QueryEscape(addrParam), nil)
w := httptest.NewRecorder()
h.handleHost(w, r)
return decodeResponse(t, w)
}
func decodeResponse(t *testing.T, w *httptest.ResponseRecorder) (int, map[string]any) {
t.Helper()
assert.Equal(t, "application/json", w.Header().Get("Content-Type"))
var body map[string]any
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body))
return w.Code, body
}
func TestInfoAPIServer_handleHost(t *testing.T) {
h, _ := newTestInfoAPIServer(t)
h.pki = newTestPKI(t, "self", []netip.Addr{netip.MustParseAddr("10.0.0.1")})
peerV4 := netip.MustParseAddr("10.0.0.99")
peerV6 := netip.MustParseAddr("fd00::99")
addTestPeer(t, h.hostMap, "laptop-alice", []netip.Addr{peerV4, peerV6},
[]netip.Prefix{netip.MustParsePrefix("192.168.50.0/24")}, []string{"eng", "ssh"})
addTestPeer(t, h.hostMap, "groupless", []netip.Addr{netip.MustParseAddr("10.0.0.77")}, nil, nil)
// an established peer comes back with its full identity
code, body := getHost(t, h, "10.0.0.99")
require.Equal(t, http.StatusOK, code)
assert.Equal(t, "laptop-alice", body["name"])
assert.Equal(t, []any{"10.0.0.99", "fd00::99"}, body["vpnAddrs"])
assert.Equal(t, []any{"10.0.0.99/24", "fd00::99/64"}, body["networks"])
assert.Equal(t, []any{"192.168.50.0/24"}, body["unsafeNetworks"])
assert.Equal(t, []any{"eng", "ssh"}, body["groups"])
assert.NotEmpty(t, body["fingerprint"])
assert.Equal(t, "2", fmt.Sprintf("%v", body["certVersion"]))
assert.NotEmpty(t, body["notBefore"])
assert.NotEmpty(t, body["notAfter"])
// empty cert slices marshal as [] rather than null
code, body = getHost(t, h, "10.0.0.77")
require.Equal(t, http.StatusOK, code)
require.NotNil(t, body["groups"])
assert.Empty(t, body["groups"])
require.NotNil(t, body["unsafeNetworks"])
assert.Empty(t, body["unsafeNetworks"])
// a port in addr is ignored so RemoteAddr can be passed through directly, including the
// bracketed v6 and 4in6 forms
for _, q := range []string{"10.0.0.99:54321", "[fd00::99]:443", "::ffff:10.0.0.99"} {
code, body = getHost(t, h, q)
require.Equal(t, http.StatusOK, code, "addr=%q", q)
assert.Equal(t, "laptop-alice", body["name"], "addr=%q", q)
}
// our own address answers from the local cert state
code, body = getHost(t, h, "10.0.0.1")
require.Equal(t, http.StatusOK, code)
assert.Equal(t, "self", body["name"])
code, body = getHost(t, h, "10.0.0.42")
assert.Equal(t, http.StatusNotFound, code)
assert.NotEmpty(t, body["error"])
// a tunnel mid-teardown (no peer cert) is treated as unknown
h.hostMap.unlockedAddHostInfo(&HostInfo{
ConnectionState: &ConnectionState{},
vpnAddrs: []netip.Addr{netip.MustParseAddr("10.0.0.66")},
relayState: RelayState{
relayForByAddr: map[netip.Addr]*Relay{},
relayForByIdx: map[uint32]*Relay{},
},
}, &Interface{})
code, _ = getHost(t, h, "10.0.0.66")
assert.Equal(t, http.StatusNotFound, code)
code, body = getHost(t, h, "not-an-address")
assert.Equal(t, http.StatusBadRequest, code)
assert.NotEmpty(t, body["error"])
r := httptest.NewRequest(http.MethodGet, "/v1/host", nil)
w := httptest.NewRecorder()
h.handleHost(w, r)
code, body = decodeResponse(t, w)
assert.Equal(t, http.StatusBadRequest, code)
assert.NotEmpty(t, body["error"])
}
func TestInfoAPIServer_handleSelf(t *testing.T) {
h, _ := newTestInfoAPIServer(t)
h.pki = newTestPKI(t, "lighthouse", []netip.Addr{netip.MustParseAddr("10.0.0.1")})
r := httptest.NewRequest(http.MethodGet, "/v1/self", nil)
w := httptest.NewRecorder()
h.handleSelf(w, r)
code, body := decodeResponse(t, w)
require.Equal(t, http.StatusOK, code)
assert.Equal(t, "lighthouse", body["name"])
assert.Equal(t, []any{"10.0.0.1"}, body["vpnAddrs"])
// no cert state available should be an error, not a panic
h.pki = nil
w = httptest.NewRecorder()
h.handleSelf(w, r)
code, body = decodeResponse(t, w)
assert.Equal(t, http.StatusInternalServerError, code)
assert.NotEmpty(t, body["error"])
}
func unixHTTPClient(path string) *http.Client {
return &http.Client{
Timeout: time.Second,
Transport: &http.Transport{
DialContext: func(ctx context.Context, _, _ string) (net.Conn, error) {
return (&net.Dialer{}).DialContext(ctx, "unix", path)
},
},
}
}
// waitForServe polls until a GET /v1/self through client succeeds
func waitForServe(t *testing.T, client *http.Client) {
t.Helper()
waitFor(t, func() bool {
resp, err := client.Get("http://hostquery/v1/self")
if err != nil {
return false
}
resp.Body.Close()
return resp.StatusCode == http.StatusOK
})
}
func skipIfNoUnixSockets(t *testing.T) {
t.Helper()
if runtime.GOOS == "windows" {
t.Skip("unix socket tests are not supported on windows CI")
}
}
func TestInfoAPIServer_unixLifecycle(t *testing.T) {
skipIfNoUnixSockets(t)
h, c := newTestInfoAPIServer(t)
h.pki = newTestPKI(t, "self", []netip.Addr{netip.MustParseAddr("10.0.0.1")})
sock := filepath.Join(t.TempDir(), "hq.sock")
setInfoAPIConfig(c, true, "unix://"+sock, "")
require.NoError(t, h.reload(c, true))
done := make(chan struct{})
go func() {
h.Start()
close(done)
}()
client := unixHTTPClient(sock)
waitForServe(t, client)
fi, err := os.Stat(sock)
require.NoError(t, err)
assert.Equal(t, fs.FileMode(0o600), fi.Mode().Perm())
resp, err := client.Get("http://hostquery/v1/host?addr=10.0.0.1")
require.NoError(t, err)
resp.Body.Close()
assert.Equal(t, http.StatusOK, resp.StatusCode)
h.Stop()
select {
case <-done:
case <-time.After(5 * time.Second):
t.Fatal("Start did not return after Stop")
}
_, err = os.Stat(sock)
assert.True(t, os.IsNotExist(err), "socket file should be unlinked on shutdown")
}
func TestInfoAPIServer_staleSocket(t *testing.T) {
skipIfNoUnixSockets(t)
h, _ := newTestInfoAPIServer(t)
sock := filepath.Join(t.TempDir(), "hq.sock")
// simulate an unclean exit, a leftover socket file with no listener
stale, err := net.ListenUnix("unix", &net.UnixAddr{Name: sock, Net: "unix"})
require.NoError(t, err)
stale.SetUnlinkOnClose(false)
require.NoError(t, stale.Close())
_, err = os.Stat(sock)
require.NoError(t, err, "stale socket file should exist")
cfg := infoAPIConfig{addr: sock, socketMode: 0o600}
ln, err := h.listen(cfg)
require.NoError(t, err, "a stale socket should be removed and rebound")
require.NoError(t, ln.Close())
}
func TestInfoAPIServer_existingFileNotReplaced(t *testing.T) {
skipIfNoUnixSockets(t)
h, _ := newTestInfoAPIServer(t)
path := filepath.Join(t.TempDir(), "hq.sock")
require.NoError(t, os.WriteFile(path, []byte("precious"), 0o600))
cfg := infoAPIConfig{addr: path, socketMode: 0o600}
_, err := h.listen(cfg)
require.Error(t, err, "a non-socket file at the listen path must not be replaced")
content, err := os.ReadFile(path)
require.NoError(t, err)
assert.Equal(t, "precious", string(content))
}
func TestInfoAPIServer_reload(t *testing.T) {
skipIfNoUnixSockets(t)
h, c := newTestInfoAPIServer(t)
h.pki = newTestPKI(t, "self", []netip.Addr{netip.MustParseAddr("10.0.0.1")})
dir := t.TempDir()
sock1 := filepath.Join(dir, "hq1.sock")
sock2 := filepath.Join(dir, "hq2.sock")
// initial reload only records config, Control.Start is what launches the runtime
setInfoAPIConfig(c, false, "unix://"+sock1, "")
require.NoError(t, h.reload(c, true))
assert.False(t, h.enabled.Load())
h.runMu.Lock()
assert.Nil(t, h.run)
h.runMu.Unlock()
// enabling via reload spawns the listener
setInfoAPIConfig(c, true, "unix://"+sock1, "")
require.NoError(t, h.reload(c, false))
waitForServe(t, unixHTTPClient(sock1))
// changing the listen path restarts on the new address
setInfoAPIConfig(c, true, "unix://"+sock2, "")
require.NoError(t, h.reload(c, false))
waitForServe(t, unixHTTPClient(sock2))
waitFor(t, func() bool {
_, err := os.Stat(sock1)
return os.IsNotExist(err)
})
// reloading an unchanged config does not restart the runtime
h.runMu.Lock()
rt := h.run
h.runMu.Unlock()
require.NoError(t, h.reload(c, false))
h.runMu.Lock()
assert.Same(t, rt, h.run)
h.runMu.Unlock()
// disabling stops the listener
setInfoAPIConfig(c, false, "unix://"+sock2, "")
require.NoError(t, h.reload(c, false))
assert.False(t, h.enabled.Load())
waitFor(t, func() bool {
h.runMu.Lock()
defer h.runMu.Unlock()
return h.run == nil
})
}
+4 -10
View File
@@ -297,7 +297,7 @@ func (f *Interface) SendVia(via *HostInfo,
c := via.ConnectionState.messageCounter.Add(1)
out = header.Encode(out, header.Version, header.Message, header.MessageRelay, relay.RemoteIndex, c)
f.connectionManager.Out(via)
f.connectionManager.OutRelay(via)
// Authenticate the header and payload, but do not encrypt for this message type.
// The payload consists of the inner, unencrypted Nebula header, as well as the end-to-end encrypted payload.
@@ -365,17 +365,11 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
//l.WithField("trace", string(debug.Stack())).Error("out Header ", &Header{Version, t, st, 0, hostinfo.remoteIndexId, c}, p)
out = header.Encode(out, header.Version, t, st, hostinfo.remoteIndexId, c)
f.connectionManager.Out(hostinfo)
// Query our LH if we haven't since the last time we've been rebound, this will cause the remote to punch against
// all our addrs and enable a faster roaming.
if t != header.CloseTunnel && hostinfo.lastRebindCount != f.rebindCount {
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
// We rebound since this tunnel last sent, ask the lighthouse to get the far side punching at us again
if f.connectionManager.Out(hostinfo) && t != header.CloseTunnel {
f.lightHouse.QueryServer(hostinfo.vpnAddrs[0])
hostinfo.lastRebindCount = f.rebindCount
if f.l.Enabled(context.Background(), slog.LevelDebug) {
f.l.Debug("Lighthouse update triggered for punch due to rebind counter",
f.l.Debug("Lighthouse update triggered for punch due to rebind epoch",
"vpnAddrs", hostinfo.vpnAddrs,
)
}
+2 -2
View File
@@ -82,8 +82,8 @@ type Interface struct {
sendRecvErrorConfig recvErrorConfig
acceptRecvErrorConfig recvErrorConfig
// rebindCount is used to decide if an active tunnel should trigger a punch notification through a lighthouse
rebindCount int8
// Bumped on every udp rebind, tunnels compare it to decide they need a punch from the far side
rebindEpoch atomic.Uint32
version string
conntrackCacheTimeout time.Duration
-3
View File
@@ -260,8 +260,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
return nil, util.ContextualizeIfNeeded("Failed to start stats emitter", err)
}
infoAPI := newInfoAPIServerFromConfig(ctx, l, pki, hostMap, c)
if configTest {
return nil, nil
}
@@ -281,7 +279,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
sshStart: sshStart,
statsStart: stats.Start,
dnsStart: ds.Start,
infoAPIStart: infoAPI.Start,
lighthouseStart: lightHouse.StartUpdateWorker,
networkChangeStart: networkChanges.Start,
connectionManagerStart: connManager.Start,