Compare commits

..
Author SHA1 Message Date
Matt RichardsonandClaude Opus 5 14a4d87faf Add nebula ctl, a local socket for the debug commands
Every diagnostic command nebula has was reachable through exactly one door:
the built-in ssh debug server. That server is off by default, and turning it
on means generating a host key, writing an sshd block with authorized public
keys, and SIGHUPing the daemon. That is a lot of ceremony to answer "what
version is this node running".

Nebula now serves the same commands over a local unix socket, enabled by
default, and `nebula ctl <command>` runs them. The socket lives in a 0700
directory so filesystem permissions are the access control; no keys, nothing
on the network. Failing to create it is logged and never blocks startup.

The command registry was already transport neutral, so this is mostly new
transport rather than new commands:

  - diag/ holds the registry, dispatch, writer and wire protocol, moved out
    of sshd because none of it was ever about ssh. sshd and ctl.go dispatch
    against one shared registry.
  - commands.go holds every command implementation, moved out of ssh.go
    (which was 85% not ssh) and renamed off the ssh prefix. Adding a command
    there makes it available over both transports.
  - ssh.go keeps only host keys, authorized users, and the listen address.
  - ctl.go supervises the socket, following the statsServer lifecycle shape.

The wire protocol frames the response rather than terminating it, because
print-cert -raw and list-hostmap -json both emit arbitrary bytes that no
sentinel could safely delimit. argv travels as a list so quoting survives.
Exit statuses are real: 0, 2 for usage, 127 for an unknown command.

Two things fall out. The ssh console now reports a real exit status instead
of a hardcoded zero, so `ssh host list-hostmap` is scriptable too. And eight
command callbacks that silently returned nil on a flags type mismatch now
report it, which the exit status makes visible.

Windows is a stub returning a clear "not supported" until it gets a named
pipe with a security descriptor; iOS and Android are never enabled, having no
daemon for a CLI to attach to.

Breaking for embedders of the sshd package: NewSSHServer takes a
*diag.Registry, SSHServer.RegisterCommand is gone in favor of registering on
that registry, and the command types live in diag rather than sshd.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014fya5fTXGiwX72FUmoL9y3
2026-09-09 17:07:35 -04:00
Nate Brown 89178f45ba windows: Fix NLMTU for ipv6 (#1871)
smoke-extra / freebsd-amd64 (push) Failing after 16s
smoke-extra / linux-amd64-ipv6disable (push) Failing after 15s
smoke-extra / netbsd-amd64 (push) Failing after 14s
smoke-extra / openbsd-amd64 (push) Failing after 17s
smoke-extra / linux-386 (push) Failing after 16s
smoke / Run multi node smoke test (push) Failing after 1m33s
Build and test / Static checks (push) Successful in 21s
Build and test / Test linux (push) Failing after 59s
Build and test / Test linux-pkcs11 (push) Failing after 1m49s
Build and test / Test linux-boringcrypto (push) Failing after 2m38s
Build and test / Test linux-fips140 (push) Failing after 2m52s
Build and test / Cross-build linux-arm (push) Successful in 2m54s
Build and test / Cross-build linux-mips (push) Successful in 3m35s
Build and test / Cross-build linux-other (push) Successful in 2m58s
Build and test / Cross-build windows (push) Successful in 59s
Build and test / Cross-build freebsd (push) Successful in 1m28s
Build and test / Cross-build netbsd (push) Successful in 1m28s
Build and test / Cross-build openbsd (push) Successful in 1m28s
Build and test / Cross-build mobile (push) Successful in 3m7s
smoke-extra / Run windows smoke test (push) Canceled after 0s
smoke / Run self traffic smoke test on macOS (push) Canceled after 0s
Build and test / Test macos (push) Canceled after 0s
Build and test / Test windows (push) Canceled after 0s
Build and test / CI status (push) Canceled after 0s
2026-09-08 18:30:16 -05:00
Wade Simmons 6a72e1c304 v1.11.1 Changelog (#1857)
smoke-extra / freebsd-amd64 (push) Failing after 18s
smoke-extra / linux-amd64-ipv6disable (push) Failing after 17s
smoke-extra / netbsd-amd64 (push) Failing after 14s
smoke-extra / openbsd-amd64 (push) Failing after 16s
smoke-extra / linux-386 (push) Failing after 17s
smoke / Run multi node smoke test (push) Failing after 1m33s
Build and test / Static checks (push) Successful in 23s
Build and test / Test linux (push) Failing after 58s
Build and test / Test linux-pkcs11 (push) Failing after 1m50s
Build and test / Test linux-boringcrypto (push) Failing after 2m36s
Build and test / Test linux-fips140 (push) Failing after 3m0s
Build and test / Cross-build linux-arm (push) Successful in 2m50s
Build and test / Cross-build linux-mips (push) Successful in 3m34s
Build and test / Cross-build linux-other (push) Successful in 2m59s
Build and test / Cross-build windows (push) Successful in 1m0s
Build and test / Cross-build freebsd (push) Successful in 1m28s
Build and test / Cross-build netbsd (push) Successful in 1m28s
Build and test / Cross-build openbsd (push) Successful in 1m28s
Build and test / Cross-build mobile (push) Successful in 3m9s
smoke-extra / Run windows smoke test (push) Canceled after 0s
smoke / Run self traffic smoke test on macOS (push) Canceled after 0s
Build and test / Test macos (push) Canceled after 0s
Build and test / Test windows (push) Canceled after 0s
Build and test / CI status (push) Canceled after 0s
Copy the v1.11.1 CHANGELOG updates from release-v1.11
2026-09-08 13:50:27 -04:00
Matt Richardson e50f8128f4 Drop hostQueries to our overlay addresses (#1866)
smoke-extra / freebsd-amd64 (push) Failing after 24s
smoke-extra / linux-amd64-ipv6disable (push) Failing after 23s
smoke-extra / netbsd-amd64 (push) Failing after 16s
smoke-extra / openbsd-amd64 (push) Failing after 11s
smoke-extra / linux-386 (push) Failing after 22s
smoke / Run multi node smoke test (push) Failing after 1m37s
Build and test / Static checks (push) Successful in 2m16s
Build and test / Test linux (push) Failing after 1m20s
Build and test / Test linux-pkcs11 (push) Failing after 1m50s
Build and test / Test linux-boringcrypto (push) Failing after 2m39s
Build and test / Test linux-fips140 (push) Failing after 2m51s
Build and test / Cross-build linux-arm (push) Successful in 2m55s
Build and test / Cross-build linux-mips (push) Successful in 3m41s
Build and test / Cross-build linux-other (push) Successful in 3m0s
Build and test / Cross-build windows (push) Successful in 1m1s
Build and test / Cross-build freebsd (push) Successful in 1m30s
Build and test / Cross-build netbsd (push) Successful in 1m29s
Build and test / Cross-build openbsd (push) Successful in 1m30s
Build and test / Cross-build mobile (push) Successful in 3m10s
smoke-extra / Run windows smoke test (push) Canceled after 0s
smoke / Run self traffic smoke test on macOS (push) Canceled after 0s
Build and test / Test macos (push) Canceled after 0s
Build and test / Test windows (push) Canceled after 0s
Build and test / CI status (push) Canceled after 0s
* Drop hostQueries to our overlay addresses

Don't bother responding to hostQuery's for our overlay addresses.
2026-09-04 11:26:08 -04:00
Jack Doan dd8f660c0a save about 10MB of RAM by not importing gopacket except for tests (#1864)
smoke-extra / freebsd-amd64 (push) Failing after 15s
smoke-extra / linux-amd64-ipv6disable (push) Failing after 15s
smoke-extra / netbsd-amd64 (push) Failing after 15s
smoke-extra / openbsd-amd64 (push) Failing after 16s
smoke-extra / linux-386 (push) Failing after 16s
smoke / Run multi node smoke test (push) Failing after 1m34s
Build and test / Static checks (push) Successful in 46s
Build and test / Test linux (push) Failing after 1m11s
Build and test / Test linux-pkcs11 (push) Failing after 1m51s
Build and test / Test linux-boringcrypto (push) Failing after 2m35s
Build and test / Test linux-fips140 (push) Failing after 2m41s
Build and test / Cross-build linux-arm (push) Successful in 2m57s
Build and test / Cross-build linux-mips (push) Successful in 4m0s
Build and test / Cross-build linux-other (push) Successful in 3m2s
Build and test / Cross-build windows (push) Successful in 1m2s
Build and test / Cross-build freebsd (push) Successful in 1m30s
Build and test / Cross-build netbsd (push) Successful in 1m29s
Build and test / Cross-build openbsd (push) Successful in 1m29s
Build and test / Cross-build mobile (push) Successful in 3m7s
smoke-extra / Run windows smoke test (push) Canceled after 0s
smoke / Run self traffic smoke test on macOS (push) Canceled after 0s
Build and test / Test macos (push) Canceled after 0s
Build and test / Test windows (push) Canceled after 0s
Build and test / CI status (push) Canceled after 0s
* save about 10MB of RAM by not importing gopacket except for tests

* big ol find-replace
2026-08-28 14:00:09 -05:00
Caleb Jasik ec3304e3a9 Recompute the transport checksum on self-forwarded packets (#1862)
smoke-extra / freebsd-amd64 (push) Failing after 16s
smoke-extra / linux-amd64-ipv6disable (push) Failing after 14s
smoke-extra / netbsd-amd64 (push) Failing after 16s
smoke-extra / openbsd-amd64 (push) Failing after 15s
smoke-extra / linux-386 (push) Failing after 16s
smoke / Run multi node smoke test (push) Failing after 1m34s
Build and test / Static checks (push) Successful in 43s
Build and test / Test linux (push) Failing after 1m29s
Build and test / Test linux-pkcs11 (push) Failing after 2m2s
Build and test / Test linux-boringcrypto (push) Failing after 2m46s
Build and test / Test linux-fips140 (push) Failing after 2m41s
Build and test / Cross-build linux-arm (push) Successful in 3m0s
Build and test / Cross-build linux-mips (push) Successful in 3m46s
Build and test / Cross-build linux-other (push) Successful in 3m8s
Build and test / Cross-build windows (push) Successful in 1m1s
Build and test / Cross-build freebsd (push) Successful in 1m33s
Build and test / Cross-build netbsd (push) Successful in 1m32s
Build and test / Cross-build openbsd (push) Successful in 1m33s
Build and test / Cross-build mobile (push) Successful in 3m16s
smoke-extra / Run windows smoke test (push) Canceled after 0s
smoke / Run self traffic smoke test on macOS (push) Canceled after 0s
Build and test / Test macos (push) Canceled after 0s
Build and test / Test windows (push) Canceled after 0s
Build and test / CI status (push) Canceled after 0s
2026-08-27 16:10:32 -05:00
Jack Doan aaa2ff7fff respect setting for tun.pin_threads_key (#1861) 2026-08-27 12:31:08 -05:00
Nate Brown 657f6ad044 Fold the rebind counter and traffic flags into one atomic word (#1820)
smoke-extra / freebsd-amd64 (push) Failing after 23s
smoke-extra / linux-amd64-ipv6disable (push) Failing after 14s
smoke-extra / netbsd-amd64 (push) Failing after 13s
smoke-extra / openbsd-amd64 (push) Failing after 13s
smoke-extra / linux-386 (push) Failing after 11s
smoke / Run multi node smoke test (push) Failing after 1m36s
Build and test / Static checks (push) Successful in 2m15s
Build and test / Test linux (push) Failing after 58s
Build and test / Test linux-pkcs11 (push) Failing after 1m55s
Build and test / Test linux-boringcrypto (push) Failing after 2m40s
Build and test / Test linux-fips140 (push) Failing after 2m41s
Build and test / Cross-build linux-arm (push) Successful in 3m3s
Build and test / Cross-build linux-mips (push) Successful in 3m43s
Build and test / Cross-build linux-other (push) Successful in 3m5s
Build and test / Cross-build windows (push) Successful in 1m2s
Build and test / Cross-build freebsd (push) Successful in 1m32s
Build and test / Cross-build netbsd (push) Successful in 1m32s
Build and test / Cross-build openbsd (push) Successful in 1m33s
Build and test / Cross-build mobile (push) Successful in 3m18s
smoke-extra / Run windows smoke test (push) Canceled after 0s
Build and test / Test macos (push) Canceled after 0s
Build and test / Test windows (push) Canceled after 0s
Build and test / CI status (push) Canceled after 0s
2026-08-26 15:15:37 -05:00
54 changed files with 4496 additions and 1216 deletions
+21
View File
@@ -41,3 +41,24 @@ jobs:
run: make fips140-all GOALS=smoke-docker
timeout-minutes: 10
smoke-self:
name: Run self traffic smoke test on macOS
runs-on: macos-latest
steps:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
with:
go-version: '1.26'
check-latest: true
- name: build
run: make bin
- name: run smoke-self
working-directory: ./.github/workflows/smoke
run: ./smoke-self.sh
timeout-minutes: 10
+130
View File
@@ -0,0 +1,130 @@
#!/bin/bash
# A host must be able to reach its own overlay address. Where the kernel sends
# that traffic through the tun rather than over loopback, nebula sees it and
# hands it straight back (immediatelyForwardToSelf), and whether the kernel
# accepts what comes back is only answerable against a real kernel. Runs one
# nebula on this machine as root and aims every probe at its own address.
set -e -x
set -o pipefail
V4=192.0.2.1
V6=2001:db8::1
case "$(uname -s)" in
Darwin) TUN_DEV=utun ;;
*) TUN_DEV=tun0 ;;
esac
ROOT="$(cd ../../.. && pwd)"
rm -rf build/self
mkdir -p build/self
cd build/self
cleanup() {
echo
echo " *** cleanup"
echo
set +e
if [ -n "$NEBULA_PID" ]
then
sudo kill "$NEBULA_PID"
fi
{ kill $(jobs -p); wait; } 2>/dev/null
sed 's/^/ [self] /' nebula.log
}
trap cleanup EXIT
# perl is on every platform this runs on; timeout(1) is not.
alarm() {
perl -e 'alarm shift; exec @ARGV' "$@"
}
RESULTS=""
FAILED=""
probe() {
local name="$1"
shift
if "$@"
then
RESULTS="$RESULTS $name=ok"
else
RESULTS="$RESULTS $name=FAIL"
FAILED="$FAILED $name"
fi
}
# Send one datagram, then wait for the listener to have written it out.
udp_probe() {
echo self | alarm 5 nc -u -w1 "$1" 3000 || true
set +x
for _ in $(seq 1 20)
do
if grep -q self "$2"
then
set -x
return 0
fi
sleep 0.25
done
set -x
return 1
}
"$ROOT/nebula-cert" ca -name "Smoke Test"
"$ROOT/nebula-cert" sign -name self -networks "$V4/24,$V6/64"
HOST=self AM_LIGHTHOUSE=true TUN_DEV="$TUN_DEV" ../../genconfig.sh >self.yml
"$ROOT/nebula" -config self.yml -test
sudo -v
sudo "$ROOT/nebula" -config self.yml >nebula.log 2>&1 &
NEBULA_PID=$!
for _ in $(seq 1 40)
do
ifconfig | grep "inet6 $V6 " >/dev/null && break
sleep 0.25
done
ifconfig | grep "inet $V4 "
ifconfig | grep "inet6 $V6 "
nc -l "$V4" 2000 >/dev/null &
nc -l "$V6" 2000 >/dev/null &
nc -u -l "$V4" 3000 >udp4.txt &
nc -u -l "$V6" 3000 >udp6.txt &
sleep 1
set +x
echo
echo " *** Testing self traffic from $V4"
echo
set -x
probe icmp4 alarm 5 ping -c1 "$V4"
probe tcp4 alarm 5 nc -z "$V4" 2000
probe udp4 udp_probe "$V4" udp4.txt
set +x
echo
echo " *** Testing self traffic from $V6"
echo
set -x
probe icmp6 alarm 5 ping6 -c1 "$V6"
probe tcp6 alarm 5 nc -z "$V6" 2000
probe udp6 udp_probe "$V6" udp6.txt
set +x
echo
echo " *** self traffic:$RESULTS"
echo
if [ -n "$FAILED" ]
then
echo "self traffic failed:$FAILED" >&2
exit 1
fi
+34 -5
View File
@@ -51,15 +51,19 @@ wsl -d $Distro -- bash -c "rm -rf $WslDir && mkdir -p $WslDir" | Out-Null
$DevName = 'nebula-smoke'
$Ip1 = '192.168.241.1'
$Ip2 = '192.168.241.2'
# Dual stack on purpose: a v4-only overlay never exercises the v6 side of tun.mtu.
$Ip6_1 = 'fd42:4242:241::1'
$Ip6_2 = 'fd42:4242:241::2'
$Mtu = 1300
$Port = 4242
& $NebulaCert ca -name 'smoke-ca' -out-crt "$WorkDir\ca.crt" -out-key "$WorkDir\ca.key"
if ($LASTEXITCODE -ne 0) { throw "nebula-cert ca failed (exit $LASTEXITCODE)" }
& $NebulaCert sign -name 'lighthouse' -networks "$Ip1/24" -ca-crt "$WorkDir\ca.crt" -ca-key "$WorkDir\ca.key" -out-crt "$WorkDir\lighthouse.crt" -out-key "$WorkDir\lighthouse.key"
& $NebulaCert sign -name 'lighthouse' -networks "$Ip1/24,$Ip6_1/64" -ca-crt "$WorkDir\ca.crt" -ca-key "$WorkDir\ca.key" -out-crt "$WorkDir\lighthouse.crt" -out-key "$WorkDir\lighthouse.key"
if ($LASTEXITCODE -ne 0) { throw "nebula-cert sign lighthouse failed (exit $LASTEXITCODE)" }
& $NebulaCert sign -name 'peer' -networks "$Ip2/24" -ca-crt "$WorkDir\ca.crt" -ca-key "$WorkDir\ca.key" -out-crt "$WorkDir\peer.crt" -out-key "$WorkDir\peer.key"
& $NebulaCert sign -name 'peer' -networks "$Ip2/24,$Ip6_2/64" -ca-crt "$WorkDir\ca.crt" -ca-key "$WorkDir\ca.key" -out-crt "$WorkDir\peer.crt" -out-key "$WorkDir\peer.key"
if ($LASTEXITCODE -ne 0) { throw "nebula-cert sign peer failed (exit $LASTEXITCODE)" }
# Windows lighthouse config.
@@ -82,7 +86,7 @@ tun:
drop_local_broadcast: false
drop_multicast: false
tx_queue: 500
mtu: 1300
mtu: $Mtu
network_category: private
logging:
level: info
@@ -126,7 +130,7 @@ tun:
drop_local_broadcast: false
drop_multicast: false
tx_queue: 500
mtu: 1300
mtu: $Mtu
logging:
level: info
format: text
@@ -169,7 +173,7 @@ Write-Host '=== WSL diagnostic ==='
wsl --version 2>&1 | Out-Host
wsl --list --verbose 2>&1 | Out-Host
wsl -d $Distro -u root -- uname -a | Out-Host
wsl -d $Distro -u root -- bash -c "modprobe tun 2>&1 || true; mkdir -p /dev/net; [ -c /dev/net/tun ] || mknod /dev/net/tun c 10 200; chmod 600 /dev/net/tun; ls -l /dev/net/tun"
wsl -d $Distro -u root -- bash -c "modprobe tun 2>&1 || true; mkdir -p /dev/net; [ -c /dev/net/tun ] || mknod /dev/net/tun c 10 200; chmod 600 /dev/net/tun; { echo 0 > /proc/sys/net/ipv6/conf/all/disable_ipv6; echo 0 > /proc/sys/net/ipv6/conf/default/disable_ipv6; } 2>/dev/null || true; ls -l /dev/net/tun"
if ($LASTEXITCODE -ne 0) { throw "failed to prepare /dev/net/tun in WSL (TUN support missing?)" }
# Deliberately no New-NetFirewallRule calls here -- nebula's windows_bypass_wdf
@@ -214,6 +218,16 @@ try {
}
Write-Host "OK: $DevName NetworkCategory=Private"
# v6 silently kept the adapter default of 65535 while v4 was correct.
foreach ($family in @('IPv4', 'IPv6')) {
Wait-Until -TimeoutSec 30 -What "$DevName $family NlMtu=$Mtu" -Predicate {
if ($lhProc.HasExited) { throw "lighthouse exited (code $($lhProc.ExitCode)) before $family mtu was set" }
$rows = @(Get-NetIPInterface -InterfaceAlias $DevName -AddressFamily $family -ErrorAction SilentlyContinue)
$rows.Count -gt 0 -and -not ($rows | Where-Object { $_.NlMtu -ne $Mtu })
}
Write-Host "OK: $DevName $family NlMtu=$Mtu"
}
Wait-Until -TimeoutSec 30 -What "WSL nebula1 with $Ip2" -Predicate {
if ($peerProc.HasExited) { throw "peer exited (code $($peerProc.ExitCode)) before tun was ready" }
$r = wsl -d $Distro -u root -- bash -c "ip -o addr show nebula1 2>/dev/null | grep -q 'inet $Ip2' && echo yes"
@@ -221,6 +235,13 @@ try {
}
Write-Host "OK: WSL nebula1 has $Ip2"
Wait-Until -TimeoutSec 30 -What "WSL nebula1 with $Ip6_2" -Predicate {
if ($peerProc.HasExited) { throw "peer exited (code $($peerProc.ExitCode)) before the v6 address was up" }
$r = wsl -d $Distro -u root -- bash -c "ip -o addr show nebula1 2>/dev/null | grep -q 'inet6 $Ip6_2' && echo yes"
("$r").Trim() -eq 'yes'
}
Write-Host "OK: WSL nebula1 has $Ip6_2"
Wait-Until -TimeoutSec 30 -What "ping from WSL peer to windows lighthouse ($Ip1)" -Predicate {
if ($peerProc.HasExited) { throw "peer exited (code $($peerProc.ExitCode)) before ping succeeded" }
$r = wsl -d $Distro -u root -- bash -c "ping -c1 -W1 $Ip1 >/dev/null 2>&1 && echo OK"
@@ -234,6 +255,14 @@ try {
}
Write-Host "OK: windows lighthouse -> WSL peer"
# Otherwise the v6 networks only prove the interface exists, not that it forwards.
Wait-Until -TimeoutSec 30 -What "v6 ping from WSL peer to windows lighthouse ($Ip6_1)" -Predicate {
if ($peerProc.HasExited) { throw "peer exited (code $($peerProc.ExitCode)) before the v6 ping succeeded" }
$r = wsl -d $Distro -u root -- bash -c "ping -6 -c1 -W1 $Ip6_1 >/dev/null 2>&1 && echo OK"
("$r").Trim() -eq 'OK'
}
Write-Host "OK: WSL peer -> windows lighthouse over v6"
Write-Host ''
Write-Host 'All smoke checks passed.'
}
+36 -1
View File
@@ -7,6 +7,32 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [Unreleased]
### Added
- New `nebula ctl <command>` subcommand, which runs any of the debug and administrative commands the sshd
block exposes without requiring an ssh server, a host key, or authorized keys. Nebula serves them over a
local unix socket, configured by the new `ctl` block and enabled by default at `/run/nebula/ctl.sock` on
Linux and `/var/run/nebula/ctl.sock` elsewhere. The socket lives in a `0700` directory so filesystem
permissions are the access control; failing to create it is logged and never prevents nebula from
starting. Packagers running nebula under systemd will want `RuntimeDirectory=nebula` in the unit so the
directory exists with the right ownership. Not supported on Windows yet, and never enabled on iOS or
Android. Reloadable.
### Changed
- The ssh console now reports a real exit status for `ssh <host> <command>` rather than always reporting
success, so commands run that way are scriptable.
- The debug and administrative commands moved out of `ssh.go` into `commands.go` and are no longer tied to
ssh: both the ssh console and `nebula ctl` dispatch against one shared registry, so a command added in
one place is available over both. Embedders of the `sshd` package are affected: `sshd.NewSSHServer` now
takes a `*diag.Registry`, `sshd.SSHServer.RegisterCommand` is gone in favor of registering on that
registry directly, and the command types now live in the `diag` package rather than being re-exported
from `sshd`.
## [1.11.1] - 2026-08-21
See the [v1.11.1](https://github.com/slackhq/nebula/milestone/30?closed=1) milestone for a complete list of changes.
### Changed
- IPv6 packets whose next header is a protocol Nebula does not parse (SCTP, GRE, IP-in-IP, etc.) are now
@@ -15,11 +41,18 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
their true protocol, so only a `proto: any` rule allows them. If you carry one of these protocols over the
overlay, confirm a `proto: any` rule covers it before upgrading, it may have been passing only through this
bypass. (#1840)
- Drop the dependency on `github.com/cyberdelia/go-metrics-graphite`, which has been unmaintained for over ten
years, by inlining the small amount of code Nebula used. (#1832)
### Fixed
- The ICMPv6 type was read from the wrong byte when classifying IPv6 packets, so the echo identifier used
for conntrack was never picked up. (#1840)
- Enforce outbound message counter limits so a tunnel is rehandshaked before the counter can wrap, preventing
nonce reuse. This is unreachable in practice, but is enforced as a defense-in-depth measure. (#1841)
- Prevent `nebula-cert ca` from running out of memory on 32bit systems when generating encrypted private keys. (#1834)
- Tolerate `ErrDumpInterrupted` when listing tun addresses on Linux, so a transient interrupted netlink dump
no longer aborts startup. (#1835)
## [1.11.0] - 2026-07-23
@@ -884,7 +917,9 @@ created.)
- Initial public release.
[Unreleased]: https://github.com/slackhq/nebula/compare/v1.10.3...HEAD
[Unreleased]: https://github.com/slackhq/nebula/compare/v1.11.1...HEAD
[1.11.1]: https://github.com/slackhq/nebula/releases/tag/v1.11.1
[1.11.0]: https://github.com/slackhq/nebula/releases/tag/v1.11.0
[1.10.3]: https://github.com/slackhq/nebula/releases/tag/v1.10.3
[1.10.2]: https://github.com/slackhq/nebula/releases/tag/v1.10.2
[1.10.1]: https://github.com/slackhq/nebula/releases/tag/v1.10.1
+4 -1
View File
@@ -338,10 +338,13 @@ smoke-relay-docker: bin-docker
smoke-docker-ipv6: export SMOKE_OVERLAY_IPV6 = 1
smoke-docker-ipv6: smoke-docker
smoke-self: bin
cd .github/workflows/smoke/ && ./smoke-self.sh
smoke-vagrant/%: bin-docker build/%/nebula
cd .github/workflows/smoke/ && ./build.sh $*
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
.FORCE:
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin bin-windows bin-windows-arm64 bin-darwin bin-freebsd bin-freebsd-arm64 bin-boringcrypto bin-fips140 bin-pkcs11 bin-docker boringcrypto build-test-mobile debug docker e2e e2ev e2evv e2evvv e2evvvv e2e-bench fips140 fips140-all $(ALL_GOFIPS140:%=fips140-%) install proto release release-linux release-freebsd release-openbsd release-netbsd release-boringcrypto release-fips140 service smoke-docker smoke-relay-docker smoke-docker-ipv6 test test-pkcs11 test-cov-html vet smoke-vagrant/%
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin bin-windows bin-windows-arm64 bin-darwin bin-freebsd bin-freebsd-arm64 bin-boringcrypto bin-fips140 bin-pkcs11 bin-docker boringcrypto build-test-mobile debug docker e2e e2ev e2evv e2evvv e2evvvv e2e-bench fips140 fips140-all $(ALL_GOFIPS140:%=fips140-%) install proto release release-linux release-freebsd release-openbsd release-netbsd release-boringcrypto release-fips140 service smoke-docker smoke-relay-docker smoke-docker-ipv6 smoke-self test test-pkcs11 test-cov-html vet smoke-vagrant/%
.DEFAULT_GOAL := bin
+131
View File
@@ -0,0 +1,131 @@
package main
import (
"errors"
"flag"
"fmt"
"io/fs"
"log/slog"
"os"
"syscall"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/diag"
)
// ctlMain implements `nebula ctl <command> [args...]`, which runs a debug command against the
// nebula already running on this host. Everything after the command name is forwarded to that
// nebula verbatim and parsed there by the same flag sets the ssh console uses, so this side
// deliberately understands as little as possible about it.
//
// Returns the process exit status.
func ctlMain(argv []string) int {
fl := flag.NewFlagSet("nebula ctl", flag.ContinueOnError)
fl.Usage = func() {
out := fl.Output()
fmt.Fprintf(out, "Usage: nebula ctl [-config path] [-socket path] <command> [arguments]\n\n")
fmt.Fprintf(out, "Runs a debug command against the running nebula on this host, over its local\n")
fmt.Fprintf(out, "control socket. Run `nebula ctl` with no command for the list of commands.\n\n")
fl.PrintDefaults()
}
socket := fl.String("socket", "", "Path to the control socket. Overrides ctl.socket from the config")
configPath := fl.String("config", "", "Path to the nebula config, read only to find ctl.socket")
// The flag package stops at the first non-flag argument, which is exactly the behaviour
// wanted here: `nebula ctl -socket /x list-hostmap -json` consumes -socket, stops at
// list-hostmap, and leaves the rest untouched for the daemon to parse.
if err := fl.Parse(argv); err != nil {
// -h is a request, not a failure.
if errors.Is(err, flag.ErrHelp) {
return diag.StatusOK
}
return diag.StatusUsage
}
path := *socket
if path == "" {
path = ctlSocketPath(*configPath)
}
if path == "" {
fmt.Fprintln(os.Stderr, "nebula ctl: no control socket path is known for this platform, set ctl.socket in the config")
return diag.StatusError
}
client, err := diag.Dial(path)
if err != nil {
fmt.Fprintln(os.Stderr, ctlDialError(path, err))
return diag.StatusError
}
defer client.Close()
args := fl.Args()
status, err := client.Run(args, os.Stdout)
if err != nil {
if errors.Is(err, diag.ErrTruncated) {
fmt.Fprintf(os.Stderr, "nebula ctl: nebula closed the connection before %s finished\n", ctlCommandName(args))
return diag.StatusError
}
fmt.Fprintf(os.Stderr, "nebula ctl: %s\n", err)
if status == diag.StatusOK {
return diag.StatusError
}
}
return status
}
// ctlSocketPath finds the socket to talk to. The platform default is the primary mechanism;
// reading the config is the refinement for someone who moved the socket. It is best effort by
// design, because config.DefaultPath resolves next to the nebula binary and a packaged install
// keeps its config somewhere else entirely, so a config we cannot find is the normal case
// rather than a failure.
func ctlSocketPath(configPath string) string {
if configPath == "" {
p, err := config.DefaultPath()
if err != nil {
return diag.DefaultSocketPath()
}
configPath = p
}
c := config.NewC(slog.New(slog.DiscardHandler))
if err := c.Load(configPath); err != nil {
return diag.DefaultSocketPath()
}
return c.GetString("ctl.socket", diag.DefaultSocketPath())
}
// ctlDialError turns a connect failure into something an operator can act on. These messages
// are the entire user experience when things are not working, so they name the path and say
// what to check.
func ctlDialError(path string, err error) string {
switch {
case errors.Is(err, diag.ErrNotSupported):
return "nebula ctl is not supported on this platform yet"
case errors.Is(err, fs.ErrNotExist):
return fmt.Sprintf("nebula ctl: no control socket at %s. Is nebula running? Is ctl.enabled set to false, or ctl.socket set to another path?", path)
case errors.Is(err, syscall.ECONNREFUSED):
return fmt.Sprintf("nebula ctl: found a stale socket at %s, nebula is not listening on it", path)
case errors.Is(err, fs.ErrPermission):
return fmt.Sprintf("nebula ctl: permission denied opening %s. nebula ctl must run as the user nebula runs as, usually root", path)
default:
return fmt.Sprintf("nebula ctl: %s", err)
}
}
// ctlCommandName names the command for an error message, for the case where there isn't one.
func ctlCommandName(args []string) string {
if len(args) == 0 {
return "the command"
}
return args[0]
}
+70
View File
@@ -0,0 +1,70 @@
package main
import (
"errors"
"io/fs"
"os"
"path/filepath"
"syscall"
"testing"
"github.com/slackhq/nebula/diag"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// The daemon parses the command's own flags, so this side must consume its own and forward
// everything from the command name onwards untouched.
func TestCtlSocketPath(t *testing.T) {
t.Run("a config naming a socket is used", func(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "config.yml")
require.NoError(t, os.WriteFile(path, []byte("ctl:\n socket: /run/somewhere/ctl.sock\n"), 0600))
assert.Equal(t, "/run/somewhere/ctl.sock", ctlSocketPath(path))
})
t.Run("a config without a ctl block falls back to the platform default", func(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "config.yml")
require.NoError(t, os.WriteFile(path, []byte("pki:\n ca: /dev/null\n"), 0600))
assert.Equal(t, diag.DefaultSocketPath(), ctlSocketPath(path))
})
// A packaged install keeps its config somewhere config.DefaultPath will never look, so a
// config we cannot read is the ordinary case and must not be fatal.
t.Run("an unreadable config falls back to the platform default", func(t *testing.T) {
assert.Equal(t, diag.DefaultSocketPath(), ctlSocketPath(filepath.Join(t.TempDir(), "nope.yml")))
})
}
func TestCtlDialError(t *testing.T) {
tests := []struct {
name string
err error
wants string
}{
{"missing socket names the path and what to check", fs.ErrNotExist, "no control socket at /x/ctl.sock. Is nebula running?"},
{"a stale socket is called stale", syscall.ECONNREFUSED, "found a stale socket at /x/ctl.sock"},
{"permission denied suggests the right user", fs.ErrPermission, "must run as the user nebula runs as"},
{"an unsupported platform says so", diag.ErrNotSupported, "not supported on this platform"},
{"anything else is reported verbatim", errors.New("something else"), "something else"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
assert.Contains(t, ctlDialError("/x/ctl.sock", tt.err), tt.wants)
})
}
t.Run("a wrapped syscall error is still recognised", func(t *testing.T) {
err := &os.SyscallError{Syscall: "connect", Err: syscall.ECONNREFUSED}
assert.Contains(t, ctlDialError("/x/ctl.sock", err), "stale socket")
})
}
func TestCtlCommandName(t *testing.T) {
assert.Equal(t, "print-cert", ctlCommandName([]string{"print-cert", "-json"}))
assert.Equal(t, "the command", ctlCommandName(nil))
}
+15
View File
@@ -32,11 +32,26 @@ func init() {
}
func main() {
// Subcommands are dispatched before flag.Parse, because flag.Parse stops at the first
// non-flag argument and everything after `ctl` has to reach the running nebula's own flag
// parser untouched. Nothing here looks at -json or a vpn address.
if len(os.Args) > 1 && os.Args[1] == "ctl" {
os.Exit(ctlMain(os.Args[2:]))
}
configPath := flag.String("config", "", "Path to either a file or directory to load configuration from")
configTest := flag.Bool("test", false, "Test the config and print the end result. Non zero exit indicates a faulty config")
printVersion := flag.Bool("version", false, "Print version")
printUsage := flag.Bool("help", false, "Print command line usage")
flag.Usage = func() {
out := flag.CommandLine.Output()
fmt.Fprintf(out, "Usage of %s:\n", os.Args[0])
flag.PrintDefaults()
fmt.Fprintf(out, "\nCommands:\n")
fmt.Fprintf(out, " ctl [command]\n\tRun a debug command against the running nebula on this host.\n\tRun `nebula ctl` on its own for the list of commands.\n")
}
flag.Parse()
if *printVersion {
+922
View File
@@ -0,0 +1,922 @@
package nebula
// The commands nebula exposes for debugging and administration. They are transport neutral:
// the ssh console in ssh.go and the `nebula ctl` socket in ctl.go both dispatch against the
// registry attachCommands fills in, and a command cannot tell which one invoked it. Adding a
// command here makes it available over both.
import (
"bytes"
"encoding/json"
"errors"
"flag"
"fmt"
"log/slog"
"maps"
"net/netip"
"os"
"path/filepath"
"runtime"
"runtime/pprof"
"sort"
"strconv"
"strings"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/diag"
"github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/logging"
)
type listHostMapFlags struct {
Json bool
Pretty bool
ByIndex bool
}
type printCertFlags struct {
Json bool
Pretty bool
Raw bool
}
type printTunnelFlags struct {
Pretty bool
}
type changeRemoteFlags struct {
Address string
}
type closeTunnelFlags struct {
LocalOnly bool
}
type createTunnelFlags struct {
Address string
}
type deviceInfoFlags struct {
Json bool
Pretty bool
}
func attachCommands(l *slog.Logger, c *config.C, reg *diag.Registry, f *Interface) {
// sandboxDir defaults to a dir in temp. The intention is that end user will
// create this dir as needed. Overriding this config value to "" allows
// writing to anywhere in the system.
defaultDir := filepath.Join(os.TempDir(), "nebula-debug")
// The key is spelled for both transports now: the profile writers are reachable over
// `nebula ctl` as well, but sshd.sandbox_dir keeps working for anyone already setting it.
sandboxDir := c.GetString("ctl.sandbox_dir", c.GetString("sshd.sandbox_dir", defaultDir))
reg.RegisterCommand(&diag.Command{
Name: "list-hostmap",
ShortDescription: "List all known previously connected hosts",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := listHostMapFlags{}
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
fl.BoolVar(&s.ByIndex, "by-index", false, "gets all hosts in the hostmap from the index table")
return fl, &s
},
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdListHostMap(f.hostMap, fs, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "list-pending-hostmap",
ShortDescription: "List all handshaking hosts",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := listHostMapFlags{}
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
fl.BoolVar(&s.ByIndex, "by-index", false, "gets all hosts in the hostmap from the index table")
return fl, &s
},
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdListHostMap(f.handshakeManager, fs, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "list-lighthouse-addrmap",
ShortDescription: "List all lighthouse map entries",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := listHostMapFlags{}
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
return fl, &s
},
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdListLighthouseMap(f.lightHouse, fs, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "reload",
ShortDescription: "Reloads configuration from disk, same as sending HUP to the process",
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdReload(c, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "start-cpu-profile",
ShortDescription: "Starts a cpu profile and write output to the provided file, ex: `cpu-profile.pb.gz`",
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdStartCpuProfile(sandboxDir, fs, a, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "stop-cpu-profile",
ShortDescription: "Stops a cpu profile and writes output to the previously provided file",
Callback: func(fs any, a []string, w diag.StringWriter) error {
pprof.StopCPUProfile()
return w.WriteLine("If a CPU profile was running it is now stopped")
},
})
reg.RegisterCommand(&diag.Command{
Name: "save-heap-profile",
ShortDescription: "Saves a heap profile to the provided path, ex: `heap-profile.pb.gz`",
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdGetHeapProfile(sandboxDir, fs, a, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "mutex-profile-fraction",
ShortDescription: "Gets or sets runtime.SetMutexProfileFraction",
Callback: cmdMutexProfileFraction,
})
reg.RegisterCommand(&diag.Command{
Name: "save-mutex-profile",
ShortDescription: "Saves a mutex profile to the provided path, ex: `mutex-profile.pb.gz`",
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdGetMutexProfile(sandboxDir, fs, a, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "log-level",
ShortDescription: "Gets or sets the current log level",
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdLogLevel(l, fs, a, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "log-format",
ShortDescription: "Gets or sets the current log format",
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdLogFormat(l, fs, a, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "version",
ShortDescription: "Prints the currently running version of nebula",
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdVersion(f, fs, a, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "device-info",
ShortDescription: "Prints information about the network device.",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := deviceInfoFlags{}
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
return fl, &s
},
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdDeviceInfo(f, fs, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "print-cert",
ShortDescription: "Prints the current certificate being used or the certificate for the provided vpn addr",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := printCertFlags{}
fl.BoolVar(&s.Json, "json", false, "outputs as json")
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
fl.BoolVar(&s.Raw, "raw", false, "raw prints the PEM encoded certificate, not compatible with -json or -pretty")
return fl, &s
},
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdPrintCert(f, fs, a, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "print-tunnel",
ShortDescription: "Prints json details about a tunnel for the provided vpn addr",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := printTunnelFlags{}
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json")
return fl, &s
},
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdPrintTunnel(f, fs, a, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "print-relays",
ShortDescription: "Prints json details about all relay info",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := printTunnelFlags{}
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json")
return fl, &s
},
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdPrintRelays(f, fs, a, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "change-remote",
ShortDescription: "Changes the remote address used in the tunnel for the provided vpn addr",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := changeRemoteFlags{}
fl.StringVar(&s.Address, "address", "", "The new remote address, ip:port")
return fl, &s
},
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdChangeRemote(f, fs, a, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "close-tunnel",
ShortDescription: "Closes a tunnel for the provided vpn addr",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := closeTunnelFlags{}
fl.BoolVar(&s.LocalOnly, "local-only", false, "Disables notifying the remote that the tunnel is shutting down")
return fl, &s
},
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdCloseTunnel(f, fs, a, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "create-tunnel",
ShortDescription: "Creates a tunnel for the provided vpn address",
Help: "The lighthouses will be queried for real addresses but you can provide one as well.",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := createTunnelFlags{}
fl.StringVar(&s.Address, "address", "", "Optionally provide a real remote address, ip:port ")
return fl, &s
},
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdCreateTunnel(f, fs, a, w)
},
})
reg.RegisterCommand(&diag.Command{
Name: "query-lighthouse",
ShortDescription: "Query the lighthouses for the provided vpn address",
Help: "This command is asynchronous. Only currently known udp addresses will be printed.",
Callback: func(fs any, a []string, w diag.StringWriter) error {
return cmdQueryLighthouse(f, fs, a, w)
},
})
}
func cmdListHostMap(hl controlHostLister, a any, w diag.StringWriter) error {
fs, ok := a.(*listHostMapFlags)
if !ok {
return fmt.Errorf("internal error: expected flags to be listHostMapFlags but was %+v", a)
}
var hm []ControlHostInfo
if fs.ByIndex {
hm = listHostMapIndexes(hl)
} else {
hm = listHostMapHosts(hl)
}
sort.Slice(hm, func(i, j int) bool {
return hm[i].VpnAddrs[0].Compare(hm[j].VpnAddrs[0]) < 0
})
if fs.Json || fs.Pretty {
js := json.NewEncoder(w.GetWriter())
if fs.Pretty {
js.SetIndent("", " ")
}
err := js.Encode(hm)
if err != nil {
return nil
}
} else {
for _, v := range hm {
err := w.WriteLine(fmt.Sprintf("%s: %s", v.VpnAddrs, v.RemoteAddrs))
if err != nil {
return err
}
}
}
return nil
}
func cmdListLighthouseMap(lightHouse *LightHouse, a any, w diag.StringWriter) error {
fs, ok := a.(*listHostMapFlags)
if !ok {
return fmt.Errorf("internal error: expected flags to be listHostMapFlags but was %+v", a)
}
type lighthouseInfo struct {
VpnAddr string `json:"vpnAddr"`
Addrs *CacheMap `json:"addrs"`
}
lightHouse.RLock()
addrMap := make([]lighthouseInfo, len(lightHouse.addrMap))
x := 0
for k, v := range lightHouse.addrMap {
addrMap[x] = lighthouseInfo{
VpnAddr: k.String(),
Addrs: v.CopyCache(),
}
x++
}
lightHouse.RUnlock()
sort.Slice(addrMap, func(i, j int) bool {
return strings.Compare(addrMap[i].VpnAddr, addrMap[j].VpnAddr) < 0
})
if fs.Json || fs.Pretty {
js := json.NewEncoder(w.GetWriter())
if fs.Pretty {
js.SetIndent("", " ")
}
err := js.Encode(addrMap)
if err != nil {
return nil
}
} else {
for _, v := range addrMap {
b, err := json.Marshal(v.Addrs)
if err != nil {
return err
}
err = w.WriteLine(fmt.Sprintf("%s: %s", v.VpnAddr, string(b)))
if err != nil {
return err
}
}
}
return nil
}
// sanitizeFilePath validates that the given file path is within the sandbox directory.
// If sandboxDir is empty, the path is returned as-is for backwards compatibility.
func sanitizeFilePath(sandboxDir, filePath string) (string, error) {
if sandboxDir == "" {
return filePath, nil
}
// Clean and resolve the path relative to the sandbox directory
if !filepath.IsAbs(filePath) {
filePath = filepath.Join(sandboxDir, filePath)
}
cleaned := filepath.Clean(filePath)
// Ensure the resolved path is within the sandbox directory
cleanedSandbox := filepath.Clean(sandboxDir)
if cleaned == cleanedSandbox {
return "", fmt.Errorf("path %q resolves to the sandbox directory itself %q", filePath, sandboxDir)
}
if !strings.HasPrefix(cleaned, cleanedSandbox+string(filepath.Separator)) {
return "", fmt.Errorf("path %q is outside the sandbox directory %q", filePath, sandboxDir)
}
return cleaned, nil
}
func cmdStartCpuProfile(sandboxDir string, fs any, a []string, w diag.StringWriter) error {
if len(a) == 0 {
err := w.WriteLine("No path to write profile provided")
return err
}
filePath, err := sanitizeFilePath(sandboxDir, a[0])
if err != nil {
return w.WriteLine(err.Error())
}
file, err := os.Create(filePath)
if err != nil {
err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
return err
}
err = pprof.StartCPUProfile(file)
if err != nil {
err = w.WriteLine(fmt.Sprintf("Unable to start cpu profile: %s", err))
return err
}
err = w.WriteLine(fmt.Sprintf("Started cpu profile, issue stop-cpu-profile to write the output to %s", a))
return err
}
func cmdVersion(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
return w.WriteLine(fmt.Sprintf("%s", ifce.version))
}
func cmdQueryLighthouse(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
if len(a) == 0 {
return w.WriteLine("No vpn address was provided")
}
vpnAddr, err := netip.ParseAddr(a[0])
if err != nil {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
if !vpnAddr.IsValid() {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
var cm *CacheMap
rl := ifce.lightHouse.Query(vpnAddr)
if rl != nil {
cm = rl.CopyCache()
}
return json.NewEncoder(w.GetWriter()).Encode(cm)
}
func cmdCloseTunnel(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
flags, ok := fs.(*closeTunnelFlags)
if !ok {
return fmt.Errorf("internal error: expected flags to be closeTunnelFlags but was %+v", fs)
}
if len(a) == 0 {
return w.WriteLine("No vpn address was provided")
}
vpnAddr, err := netip.ParseAddr(a[0])
if err != nil {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
if !vpnAddr.IsValid() {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
if hostInfo == nil {
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn address: %v", a[0]))
}
if !flags.LocalOnly {
ifce.send(
header.CloseTunnel,
0,
hostInfo.ConnectionState,
hostInfo,
[]byte{},
make([]byte, 12, 12),
make([]byte, mtu),
)
}
ifce.closeTunnel(hostInfo)
return w.WriteLine("Closed")
}
func cmdCreateTunnel(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
flags, ok := fs.(*createTunnelFlags)
if !ok {
return fmt.Errorf("internal error: expected flags to be createTunnelFlags but was %+v", fs)
}
if len(a) == 0 {
return w.WriteLine("No vpn address was provided")
}
vpnAddr, err := netip.ParseAddr(a[0])
if err != nil {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
if !vpnAddr.IsValid() {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
if hostInfo != nil {
return w.WriteLine(fmt.Sprintf("Tunnel already exists"))
}
hostInfo = ifce.handshakeManager.QueryVpnAddr(vpnAddr)
if hostInfo != nil {
return w.WriteLine(fmt.Sprintf("Tunnel already handshaking"))
}
var addr netip.AddrPort
if flags.Address != "" {
addr, err = netip.ParseAddrPort(flags.Address)
if err != nil {
return w.WriteLine("Address could not be parsed")
}
}
hostInfo = ifce.handshakeManager.StartHandshake(vpnAddr, nil)
if addr.IsValid() {
hostInfo.SetRemote(addr)
}
return w.WriteLine("Created")
}
func cmdChangeRemote(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
flags, ok := fs.(*changeRemoteFlags)
if !ok {
return fmt.Errorf("internal error: expected flags to be changeRemoteFlags but was %+v", fs)
}
if len(a) == 0 {
return w.WriteLine("No vpn address was provided")
}
if flags.Address == "" {
return w.WriteLine("No address was provided")
}
addr, err := netip.ParseAddrPort(flags.Address)
if err != nil {
return w.WriteLine("Address could not be parsed")
}
vpnAddr, err := netip.ParseAddr(a[0])
if err != nil {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
if !vpnAddr.IsValid() {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
if hostInfo == nil {
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn address: %v", a[0]))
}
hostInfo.SetRemote(addr)
return w.WriteLine("Changed")
}
func cmdGetHeapProfile(sandboxDir string, fs any, a []string, w diag.StringWriter) error {
if len(a) == 0 {
return w.WriteLine("No path to write profile provided")
}
filePath, err := sanitizeFilePath(sandboxDir, a[0])
if err != nil {
return w.WriteLine(err.Error())
}
file, err := os.Create(filePath)
if err != nil {
err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
return err
}
err = pprof.WriteHeapProfile(file)
if err != nil {
err = w.WriteLine(fmt.Sprintf("Unable to write profile: %s", err))
return err
}
err = w.WriteLine(fmt.Sprintf("Mem profile created at %s", a))
return err
}
func cmdMutexProfileFraction(fs any, a []string, w diag.StringWriter) error {
if len(a) == 0 {
rate := runtime.SetMutexProfileFraction(-1)
return w.WriteLine(fmt.Sprintf("Current value: %d", rate))
}
newRate, err := strconv.Atoi(a[0])
if err != nil {
return w.WriteLine(fmt.Sprintf("Invalid argument: %s", a[0]))
}
oldRate := runtime.SetMutexProfileFraction(newRate)
return w.WriteLine(fmt.Sprintf("New value: %d. Old value: %d", newRate, oldRate))
}
func cmdGetMutexProfile(sandboxDir string, fs any, a []string, w diag.StringWriter) error {
if len(a) == 0 {
return w.WriteLine("No path to write profile provided")
}
filePath, err := sanitizeFilePath(sandboxDir, a[0])
if err != nil {
return w.WriteLine(err.Error())
}
file, err := os.Create(filePath)
if err != nil {
return w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
}
defer file.Close()
mutexProfile := pprof.Lookup("mutex")
if mutexProfile == nil {
return w.WriteLine("Unable to get pprof.Lookup(\"mutex\")")
}
err = mutexProfile.WriteTo(file, 0)
if err != nil {
return w.WriteLine(fmt.Sprintf("Unable to write profile: %s", err))
}
return w.WriteLine(fmt.Sprintf("Mutex profile created at %s", a))
}
func cmdLogLevel(l *slog.Logger, fs any, a []string, w diag.StringWriter) error {
ctrl, ok := l.Handler().(interface {
GetLevel() slog.Level
SetLevel(slog.Level)
})
if !ok {
return w.WriteLine("Log level is not reconfigurable on this logger")
}
if len(a) == 0 {
return w.WriteLine(fmt.Sprintf("Log level is: %s", logging.LevelName(ctrl.GetLevel())))
}
level, err := logging.ParseLevel(strings.ToLower(a[0]))
if err != nil {
return w.WriteLine(fmt.Sprintf("Unknown log level %s. Possible log levels: trace, debug, info, warn, error", a))
}
ctrl.SetLevel(level)
return w.WriteLine(fmt.Sprintf("Log level is: %s", logging.LevelName(ctrl.GetLevel())))
}
func cmdLogFormat(l *slog.Logger, fs any, a []string, w diag.StringWriter) error {
ctrl, ok := l.Handler().(interface {
GetFormat() string
SetFormat(string) error
})
if !ok {
return w.WriteLine("Log format is not reconfigurable on this logger")
}
if len(a) == 0 {
return w.WriteLine(fmt.Sprintf("Log format is: %s", ctrl.GetFormat()))
}
if err := ctrl.SetFormat(strings.ToLower(a[0])); err != nil {
return err
}
return w.WriteLine(fmt.Sprintf("Log format is: %s", ctrl.GetFormat()))
}
func cmdPrintCert(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
args, ok := fs.(*printCertFlags)
if !ok {
return fmt.Errorf("internal error: expected flags to be printCertFlags but was %+v", fs)
}
cert := ifce.pki.getCertState().GetDefaultCertificate()
if len(a) > 0 {
vpnAddr, err := netip.ParseAddr(a[0])
if err != nil {
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
}
if !vpnAddr.IsValid() {
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
}
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
if hostInfo == nil {
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn addr: %v", a[0]))
}
cert = hostInfo.GetCert().Certificate
}
if args.Json || args.Pretty {
b, err := cert.MarshalJSON()
if err != nil {
return nil
}
if args.Pretty {
buf := new(bytes.Buffer)
err := json.Indent(buf, b, "", " ")
b = buf.Bytes()
if err != nil {
return nil
}
}
return w.WriteBytes(b)
}
if args.Raw {
b, err := cert.MarshalPEM()
if err != nil {
return nil
}
return w.WriteBytes(b)
}
return w.WriteLine(cert.String())
}
func cmdPrintRelays(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
args, ok := fs.(*printTunnelFlags)
if !ok {
return fmt.Errorf("internal error: expected flags to be printTunnelFlags but was %+v", fs)
}
relays := map[uint32]*HostInfo{}
ifce.hostMap.Lock()
maps.Copy(relays, ifce.hostMap.Relays)
ifce.hostMap.Unlock()
type RelayFor struct {
Error error
Type string
State string
PeerAddr netip.Addr
LocalIndex uint32
RemoteIndex uint32
RelayedThrough []netip.Addr
}
type RelayOutput struct {
NebulaAddr netip.Addr
RelayForAddrs []RelayFor
}
type CmdOutput struct {
Relays []*RelayOutput
}
co := CmdOutput{}
enc := json.NewEncoder(w.GetWriter())
if args.Pretty {
enc.SetIndent("", " ")
}
for k, v := range relays {
ro := RelayOutput{NebulaAddr: v.vpnAddrs[0]}
co.Relays = append(co.Relays, &ro)
relayHI := ifce.hostMap.QueryVpnAddr(v.vpnAddrs[0])
if relayHI == nil {
ro.RelayForAddrs = append(ro.RelayForAddrs, RelayFor{Error: errors.New("could not find hostinfo")})
continue
}
for _, vpnAddr := range relayHI.relayState.CopyRelayForIps() {
rf := RelayFor{Error: nil}
r, ok := relayHI.relayState.GetRelayForByAddr(vpnAddr)
if ok {
t := ""
switch r.Type {
case ForwardingType:
t = "forwarding"
case TerminalType:
t = "terminal"
default:
t = "unknown"
}
s := ""
switch r.State {
case Requested:
s = "requested"
case Established:
s = "established"
default:
s = "unknown"
}
rf.LocalIndex = r.LocalIndex
rf.RemoteIndex = r.RemoteIndex
rf.PeerAddr = r.PeerAddr
rf.Type = t
rf.State = s
if rf.LocalIndex != k {
rf.Error = fmt.Errorf("hostmap LocalIndex '%v' does not match RelayState LocalIndex", k)
}
}
relayedHI := ifce.hostMap.QueryVpnAddr(vpnAddr)
if relayedHI != nil {
rf.RelayedThrough = append(rf.RelayedThrough, relayedHI.relayState.CopyRelayIps()...)
}
ro.RelayForAddrs = append(ro.RelayForAddrs, rf)
}
}
err := enc.Encode(co)
if err != nil {
return err
}
return nil
}
func cmdPrintTunnel(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
args, ok := fs.(*printTunnelFlags)
if !ok {
return fmt.Errorf("internal error: expected flags to be printTunnelFlags but was %+v", fs)
}
if len(a) == 0 {
return w.WriteLine("No vpn address was provided")
}
vpnAddr, err := netip.ParseAddr(a[0])
if err != nil {
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
}
if !vpnAddr.IsValid() {
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
}
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
if hostInfo == nil {
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn addr: %v", a[0]))
}
enc := json.NewEncoder(w.GetWriter())
if args.Pretty {
enc.SetIndent("", " ")
}
return enc.Encode(copyHostInfo(hostInfo, ifce.hostMap.GetPreferredRanges()))
}
func cmdDeviceInfo(ifce *Interface, fs any, w diag.StringWriter) error {
data := struct {
Name string `json:"name"`
Cidr []netip.Prefix `json:"cidr"`
}{
Name: ifce.inside.Name(),
Cidr: make([]netip.Prefix, len(ifce.inside.Networks())),
}
copy(data.Cidr, ifce.inside.Networks())
flags, ok := fs.(*deviceInfoFlags)
if !ok {
return fmt.Errorf("internal error: expected flags to be deviceInfoFlags but was %+v", fs)
}
if flags.Json || flags.Pretty {
js := json.NewEncoder(w.GetWriter())
if flags.Pretty {
js.SetIndent("", " ")
}
return js.Encode(data)
} else {
return w.WriteLine(fmt.Sprintf("name=%v cidr=%v", data.Name, data.Cidr))
}
}
func cmdReload(c *config.C, w diag.StringWriter) error {
err := w.WriteLine("Reloading config")
c.ReloadConfig()
return err
}
+69
View File
@@ -0,0 +1,69 @@
package nebula
import (
"bytes"
"log/slog"
"testing"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/diag"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// attachedCommands is every command nebula exposes. The ssh console and `nebula ctl` dispatch
// against this one set, so this list is the contract for both transports.
var attachedCommands = []string{
"change-remote",
"close-tunnel",
"create-tunnel",
"device-info",
"list-hostmap",
"list-lighthouse-addrmap",
"list-pending-hostmap",
"log-format",
"log-level",
"mutex-profile-fraction",
"print-cert",
"print-relays",
"print-tunnel",
"query-lighthouse",
"reload",
"save-heap-profile",
"save-mutex-profile",
"start-cpu-profile",
"stop-cpu-profile",
"version",
}
func TestAttachCommands(t *testing.T) {
l := slog.New(slog.DiscardHandler)
reg := diag.NewRegistry()
// The callbacks capture these but do not touch them until a command runs, and this test
// only registers and asks for help.
attachCommands(l, config.NewC(l), reg, &Interface{})
t.Run("every command is registered", func(t *testing.T) {
for _, name := range attachedCommands {
assert.Equal(t, []string{name}, reg.Match(name), "%s is not registered", name)
}
})
t.Run("help is available for every command", func(t *testing.T) {
for _, name := range attachedCommands {
buf := &bytes.Buffer{}
require.NoError(t, reg.DispatchArgs([]string{"help", name}, diag.NewWriter(buf)), name)
assert.Contains(t, buf.String(), name+" - ", name)
}
})
t.Run("the command list names them all", func(t *testing.T) {
buf := &bytes.Buffer{}
require.NoError(t, reg.DispatchArgs(nil, diag.NewWriter(buf)))
for _, name := range attachedCommands {
assert.Contains(t, buf.String(), name+" - ", name)
}
})
}
+14 -8
View File
@@ -105,11 +105,18 @@ 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)
// OutNoRebind records outbound traffic without consuming the rebind epoch, for relayed sends: the direct path
// to the relay consumes the edge, the via send must not.
func (cm *connectionManager) OutNoRebind(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 +135,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
}
@@ -346,7 +352,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
@@ -369,7 +375,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"},
@@ -420,7 +426,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])
}
@@ -326,31 +326,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])
@@ -358,9 +358,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])
}
+28 -1
View File
@@ -12,6 +12,7 @@ import (
"github.com/slackhq/nebula/handshake"
"github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/test"
"github.com/slackhq/nebula/udp"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
@@ -117,7 +118,33 @@ func TestSendNoMetricsDropsExhausted(t *testing.T) {
// The crossing send is refused: it records an exhaustion drop and never reaches connectionManager.Out.
assert.Equal(t, int64(1), f.messageMetrics.txExhausted.Count())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.sentSinceCheck())
}
// TestSendNoMetricsCloseTunnelKeepsRebindEpoch pins that a closing tunnel does not consume a rebind, a later
// packet on a re-established tunnel still needs that edge to trigger the far-side punch.
func TestSendNoMetricsCloseTunnelKeepsRebindEpoch(t *testing.T) {
initR, _ := runTestHandshake(t)
ci, err := newConnectionStateFromResult(initR)
require.NoError(t, err)
f := &Interface{
l: test.NewLogger(),
messageMetrics: &MessageMetrics{txExhausted: metrics.NewCounter()},
writers: []udp.Conn{udp.NoopConn{}},
connectionManager: &connectionManager{},
}
hostinfo := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("10.0.0.1")}, ConnectionState: ci}
// Tunnel is on epoch 0, then we rebind.
hostinfo.markOut(0)
f.rebindEpoch.Add(1)
remote := netip.MustParseAddrPort("10.0.0.2:4242")
f.sendNoMetrics(header.CloseTunnel, 0, ci, hostinfo, remote, []byte{}, make([]byte, 12), make([]byte, mtu), 0)
// markOut at the new epoch still reports the move, so the edge was preserved.
assert.True(t, hostinfo.markOut(1), "a CloseTunnel send must not consume the rebind epoch")
}
func TestNewConnectionStateFromResult(t *testing.T) {
+5 -1
View File
@@ -50,6 +50,7 @@ type Control struct {
ctx context.Context
cancel context.CancelFunc
sshStart func()
ctlStart func()
statsStart func()
dnsStart func()
lighthouseStart func()
@@ -99,6 +100,9 @@ func (c *Control) Start() error {
if c.sshStart != nil {
go c.sshStart()
}
if c.ctlStart != nil {
go c.ctlStart()
}
if c.statsStart != nil {
go c.statsStart()
}
@@ -212,7 +216,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 {
+234
View File
@@ -0,0 +1,234 @@
package nebula
import (
"context"
"errors"
"log/slog"
"net"
"path/filepath"
"sync"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/diag"
"github.com/slackhq/nebula/util"
)
// ctlConfig is the parsed form of the `ctl` config block. It is comparable so that a reload
// can tell "nothing changed" from "the socket moved" with ==.
type ctlConfig struct {
enabled bool
socket string
// explicit records that the operator named a socket path rather than taking the platform
// default. It only affects how loudly a failure to listen is reported: an unprivileged
// nebula that cannot create /run/nebula is a normal deployment, not a problem to shout
// about on every upgrade, but a path someone chose deliberately failing to bind is.
explicit bool
}
// ctlServer owns the unix socket `nebula ctl` connects to. It exposes the same command
// registry the ssh console does, minus the ceremony of running an ssh server: the socket is
// local only and guarded by filesystem permissions, so it needs no keys.
//
// The lifecycle mirrors statsServer: the constructor wires the reload callback, reload
// records config and reconciles a running listener, Start builds and serves the runtime, and
// Stop tears it down.
type ctlServer struct {
l *slog.Logger
ctx context.Context
srv *diag.Server
runMu sync.Mutex
runCfg *ctlConfig
run *ctlRuntime
}
// ctlRuntime is the live state owned by a single Start invocation.
type ctlRuntime struct {
cancel context.CancelFunc
listener net.Listener
}
// newCtlServerFromConfig builds a ctlServer, parses the config, and registers a reload
// callback. It deliberately does not start listening: there is no interface yet, and
// Control.Start is what launches the first runtime. The callback is registered before the
// config is parsed so a SIGHUP can fix a bad block even if the first parse failed.
//
// reg is only held, never read, until Start runs. That is what lets this be constructed
// before attachCommands has populated the registry.
func newCtlServerFromConfig(ctx context.Context, l *slog.Logger, c *config.C, reg *diag.Registry) (*ctlServer, error) {
s := &ctlServer{
l: l,
ctx: ctx,
srv: diag.NewServer(l, reg),
}
c.RegisterReloadCallback(func(c *config.C) {
if err := s.reload(c, false); err != nil {
s.l.Error("Failed to reload ctl from config", "error", err)
}
})
if err := s.reload(c, true); err != nil {
return s, err
}
return s, nil
}
// loadCtlConfig parses and validates the `ctl` block. An empty socket path while enabled is
// not an error: it means the platform has no default and the operator did not name one, so
// there is simply nothing to listen on.
func loadCtlConfig(c *config.C) (ctlConfig, error) {
cfg := ctlConfig{
enabled: c.GetBool("ctl.enabled", true),
socket: c.GetString("ctl.socket", diag.DefaultSocketPath()),
explicit: c.IsSet("ctl.socket"),
}
if cfg.enabled && cfg.socket != "" && !filepath.IsAbs(cfg.socket) {
return cfg, util.NewContextualError("ctl.socket must be an absolute path", m{"path": cfg.socket}, nil)
}
return cfg, nil
}
// reload parses the config and records it, then reconciles the running listener against it:
//
// - newly enabled -> spawn Start
// - newly disabled -> Stop the runtime
// - socket moved (still enabled) -> Stop the old, Start the new
// - no change -> no-op
//
// On the initial call it only records configuration; Control.Start is what launches the first
// runtime via ctlStart. There is no interface to serve yet at that point.
func (s *ctlServer) reload(c *config.C, initial bool) error {
newCfg, err := loadCtlConfig(c)
if err != nil {
return err
}
s.runMu.Lock()
sameCfg := s.runCfg != nil && *s.runCfg == newCfg
s.runCfg = &newCfg
running := s.run != nil
s.runMu.Unlock()
if initial || sameCfg {
return nil
}
if running {
s.Stop()
}
if newCfg.enabled && newCfg.socket != "" {
go s.Start()
}
return nil
}
// Start binds the socket and serves until Stop is called or ctx fires. Safe to call when ctl
// is disabled or already running: both no-op.
func (s *ctlServer) Start() {
s.runMu.Lock()
if s.ctx.Err() != nil || s.run != nil || s.runCfg == nil {
s.runMu.Unlock()
return
}
cfg := *s.runCfg
s.runMu.Unlock()
if !cfg.enabled || cfg.socket == "" {
if cfg.enabled {
s.l.Info("ctl has no socket path on this platform, `nebula ctl` will not be available",
"hint", "set ctl.socket to enable it",
)
}
return
}
listener, err := diag.Listen(cfg.socket)
if err != nil {
// A default path nebula cannot create is an ordinary state for an unprivileged
// install; a path the operator chose failing to bind is something they want to know
// about. Either way ctl is optional and nebula carries on without it.
if cfg.explicit {
s.l.Error("Failed to listen on the ctl socket", "ctlSocket", cfg.socket, "error", err)
} else {
s.l.Info("Not serving the ctl socket, `nebula ctl` will not be available",
"ctlSocket", cfg.socket,
"error", err,
"hint", "set ctl.socket to a path nebula can write, or ctl.enabled to false",
)
}
// Drop the cached config so a SIGHUP retries once the underlying problem is fixed,
// even when the config itself is unchanged.
s.runMu.Lock()
if s.runCfg != nil && *s.runCfg == cfg {
s.runCfg = nil
}
s.runMu.Unlock()
return
}
runCtx, cancel := context.WithCancel(s.ctx)
rt := &ctlRuntime{cancel: cancel, listener: listener}
s.runMu.Lock()
// Losing the race against a Stop or a competing Start means this listener is already
// obsolete. Close it rather than serving a socket nobody will tear down.
if s.ctx.Err() != nil || s.run != nil {
s.runMu.Unlock()
cancel()
_ = listener.Close()
return
}
s.run = rt
s.runMu.Unlock()
s.l.Info("ctl socket is listening", "ctlSocket", cfg.socket)
err = s.srv.Serve(runCtx, listener)
if err != nil {
s.l.Error("The ctl listener stopped", "ctlSocket", cfg.socket, "error", err)
}
// Clear our runtime only if nothing has replaced it.
s.runMu.Lock()
if s.run == rt {
rt.cancel()
s.run = nil
if err != nil {
// An unclean exit leaves runCfg cached as if it were applied, so drop it and let a
// SIGHUP retry.
s.runCfg = nil
}
}
s.runMu.Unlock()
}
// Stop closes the listener and unlinks the socket. It deliberately does not touch connections
// that are already being served: `nebula ctl reload` runs every reload callback inline on its
// own connection, including this one, and hanging up on it would truncate the response to a
// reload that actually succeeded.
//
// The socket file is removed by net.UnixListener's unlink-on-close, so there is no os.Remove
// here; doing it by hand would delete a successor's socket after a fast reload.
func (s *ctlServer) Stop() {
s.runMu.Lock()
rt := s.run
s.run = nil
s.runMu.Unlock()
if rt == nil {
return
}
rt.cancel()
if err := rt.listener.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
s.l.Warn("Failed to close the ctl listener", "error", err)
}
}
+305
View File
@@ -0,0 +1,305 @@
//go:build !windows
package nebula
import (
"context"
"log/slog"
"os"
"path/filepath"
"testing"
"time"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/diag"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func newTestCtlServer(t *testing.T) (*ctlServer, *config.C) {
t.Helper()
l := slog.New(slog.DiscardHandler)
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
return &ctlServer{
l: l,
ctx: ctx,
srv: diag.NewServer(l, diag.NewRegistry()),
}, config.NewC(l)
}
func setCtlConfig(c *config.C, m map[string]any) {
c.Settings["ctl"] = m
}
func currentCtlRuntime(s *ctlServer) *ctlRuntime {
s.runMu.Lock()
defer s.runMu.Unlock()
return s.run
}
// testCtlSocket returns a short socket path, see the note in diag/server_test.go about
// sun_path on darwin.
func testCtlSocket(t *testing.T) string {
t.Helper()
dir, err := os.MkdirTemp("/tmp", "nebctl")
require.NoError(t, err)
t.Cleanup(func() { _ = os.RemoveAll(dir) })
return filepath.Join(dir, "ctl.sock")
}
func startCtl(t *testing.T, s *ctlServer) chan struct{} {
t.Helper()
done := make(chan struct{})
go func() {
s.Start()
close(done)
}()
return done
}
func requireCtlStopped(t *testing.T, done chan struct{}) {
t.Helper()
select {
case <-done:
case <-time.After(5 * time.Second):
t.Fatal("ctl Start did not return after Stop")
}
}
func TestCtlServer_loadConfig(t *testing.T) {
t.Run("defaults to enabled at the platform path", func(t *testing.T) {
_, c := newTestCtlServer(t)
cfg, err := loadCtlConfig(c)
require.NoError(t, err)
assert.True(t, cfg.enabled)
assert.Equal(t, diag.DefaultSocketPath(), cfg.socket)
assert.False(t, cfg.explicit)
})
t.Run("an operator chosen path is recorded as explicit", func(t *testing.T) {
_, c := newTestCtlServer(t)
setCtlConfig(c, map[string]any{"socket": "/run/somewhere/ctl.sock"})
cfg, err := loadCtlConfig(c)
require.NoError(t, err)
assert.Equal(t, "/run/somewhere/ctl.sock", cfg.socket)
assert.True(t, cfg.explicit)
})
t.Run("a relative path is rejected", func(t *testing.T) {
_, c := newTestCtlServer(t)
setCtlConfig(c, map[string]any{"socket": "ctl.sock"})
_, err := loadCtlConfig(c)
require.Error(t, err)
assert.Contains(t, err.Error(), "must be an absolute path")
})
t.Run("a relative path is not rejected when ctl is off", func(t *testing.T) {
_, c := newTestCtlServer(t)
setCtlConfig(c, map[string]any{"enabled": false, "socket": "ctl.sock"})
_, err := loadCtlConfig(c)
assert.NoError(t, err)
})
}
func TestCtlServer_reload(t *testing.T) {
t.Run("the initial reload records config without listening", func(t *testing.T) {
s, c := newTestCtlServer(t)
setCtlConfig(c, map[string]any{"socket": testCtlSocket(t)})
require.NoError(t, s.reload(c, true))
assert.Nil(t, currentCtlRuntime(s), "Control.Start is what starts listening")
})
t.Run("enabling on reload starts listening", func(t *testing.T) {
s, c := newTestCtlServer(t)
path := testCtlSocket(t)
setCtlConfig(c, map[string]any{"enabled": false, "socket": path})
require.NoError(t, s.reload(c, true))
setCtlConfig(c, map[string]any{"enabled": true, "socket": path})
require.NoError(t, s.reload(c, false))
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
assert.FileExists(t, path)
s.Stop()
})
t.Run("disabling on reload stops listening and unlinks", func(t *testing.T) {
s, c := newTestCtlServer(t)
path := testCtlSocket(t)
setCtlConfig(c, map[string]any{"enabled": true, "socket": path})
require.NoError(t, s.reload(c, true))
done := startCtl(t, s)
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
setCtlConfig(c, map[string]any{"enabled": false, "socket": path})
require.NoError(t, s.reload(c, false))
requireCtlStopped(t, done)
assert.Nil(t, currentCtlRuntime(s))
assert.NoFileExists(t, path)
})
t.Run("moving the socket restarts at the new path", func(t *testing.T) {
s, c := newTestCtlServer(t)
oldPath := testCtlSocket(t)
newPath := testCtlSocket(t)
setCtlConfig(c, map[string]any{"socket": oldPath})
require.NoError(t, s.reload(c, true))
done := startCtl(t, s)
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
require.FileExists(t, oldPath)
setCtlConfig(c, map[string]any{"socket": newPath})
require.NoError(t, s.reload(c, false))
requireCtlStopped(t, done)
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
assert.FileExists(t, newPath)
assert.NoFileExists(t, oldPath, "the old socket should have been unlinked")
s.Stop()
})
t.Run("an unchanged config leaves the listener alone", func(t *testing.T) {
s, c := newTestCtlServer(t)
setCtlConfig(c, map[string]any{"socket": testCtlSocket(t)})
require.NoError(t, s.reload(c, true))
startCtl(t, s)
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
before := currentCtlRuntime(s)
require.NoError(t, s.reload(c, false))
assert.Same(t, before, currentCtlRuntime(s), "the runtime should not have been replaced")
s.Stop()
})
}
func TestCtlServer_Start(t *testing.T) {
t.Run("a command can be run over the socket", func(t *testing.T) {
l := slog.New(slog.DiscardHandler)
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
reg := diag.NewRegistry()
s := &ctlServer{l: l, ctx: ctx, srv: diag.NewServer(l, reg)}
c := config.NewC(l)
path := testCtlSocket(t)
setCtlConfig(c, map[string]any{"socket": path})
require.NoError(t, s.reload(c, true))
startCtl(t, s)
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
client, err := diag.Dial(path)
require.NoError(t, err)
defer client.Close()
out := &testWriter{}
status, err := client.Run([]string{"help"}, out)
require.NoError(t, err)
assert.Equal(t, diag.StatusOK, status)
assert.Contains(t, out.String(), "Available commands:")
s.Stop()
})
t.Run("Start is a no-op when ctl is disabled", func(t *testing.T) {
s, c := newTestCtlServer(t)
setCtlConfig(c, map[string]any{"enabled": false, "socket": testCtlSocket(t)})
require.NoError(t, s.reload(c, true))
s.Start()
assert.Nil(t, currentCtlRuntime(s))
})
t.Run("Start is a no-op with no socket path for this platform", func(t *testing.T) {
s, c := newTestCtlServer(t)
setCtlConfig(c, map[string]any{"enabled": true, "socket": ""})
require.NoError(t, s.reload(c, true))
s.Start()
assert.Nil(t, currentCtlRuntime(s))
})
t.Run("Start is a no-op after the context is cancelled", func(t *testing.T) {
l := slog.New(slog.DiscardHandler)
ctx, cancel := context.WithCancel(context.Background())
s := &ctlServer{l: l, ctx: ctx, srv: diag.NewServer(l, diag.NewRegistry())}
c := config.NewC(l)
path := testCtlSocket(t)
setCtlConfig(c, map[string]any{"socket": path})
require.NoError(t, s.reload(c, true))
cancel()
s.Start()
assert.Nil(t, currentCtlRuntime(s))
assert.NoFileExists(t, path)
})
// A path nebula cannot bind must not stop it from running, and a SIGHUP with the same
// config has to be able to retry once the problem is fixed.
t.Run("a listen failure is survivable and retried on the next reload", func(t *testing.T) {
s, c := newTestCtlServer(t)
path := testCtlSocket(t)
require.NoError(t, os.MkdirAll(filepath.Dir(path), 0700))
require.NoError(t, os.WriteFile(path, []byte("in the way"), 0600))
setCtlConfig(c, map[string]any{"socket": path})
require.NoError(t, s.reload(c, true))
s.Start()
assert.Nil(t, currentCtlRuntime(s))
s.runMu.Lock()
cachedCfg := s.runCfg
s.runMu.Unlock()
assert.Nil(t, cachedCfg, "the cached config should be dropped so a reload retries")
require.NoError(t, os.Remove(path))
require.NoError(t, s.reload(c, false))
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
s.Stop()
})
t.Run("Stop is idempotent", func(t *testing.T) {
s, c := newTestCtlServer(t)
setCtlConfig(c, map[string]any{"socket": testCtlSocket(t)})
require.NoError(t, s.reload(c, true))
done := startCtl(t, s)
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
s.Stop()
requireCtlStopped(t, done)
assert.NotPanics(t, s.Stop)
})
}
// testWriter collects command output.
type testWriter struct{ b []byte }
func (w *testWriter) Write(p []byte) (int, error) {
w.b = append(w.b, p...)
return len(p), nil
}
func (w *testWriter) String() string { return string(w.b) }
+41
View File
@@ -0,0 +1,41 @@
package diag
import (
"bufio"
"io"
"net"
"time"
)
// dialTimeout bounds the connect only. A command may take as long as it likes to answer.
const dialTimeout = 2 * time.Second
// Client is a connection to a nebula serving the ctl socket. It carries exactly one command.
type Client struct {
conn net.Conn
}
// Dial connects to the nebula serving at path. On a platform without socket support the
// returned error wraps ErrNotSupported.
func Dial(path string) (*Client, error) {
conn, err := dialSocket(path, dialTimeout)
if err != nil {
return nil, err
}
return &Client{conn: conn}, nil
}
// Run sends args and streams the command's output to out, returning the command's exit
// status. A non-nil error means the exchange itself failed and the status means nothing.
func (c *Client) Run(args []string, out io.Writer) (int, error) {
if err := writeRequest(c.conn, args); err != nil {
return 0, err
}
return readResponse(bufio.NewReader(c.conn), out)
}
func (c *Client) Close() error {
return c.conn.Close()
}
+16 -3
View File
@@ -1,4 +1,4 @@
package sshd
package diag
import (
"errors"
@@ -10,6 +10,17 @@ import (
"github.com/armon/go-radix"
)
var (
// ErrUnknownCommand is returned by the Registry when the first argument names no
// registered command. The user has already been told so on their writer.
ErrUnknownCommand = errors.New("unknown command")
// ErrUsage wraps a flag parsing failure. The flag package has already written the
// details to the caller's writer by the time this is returned, so a transport should
// use it only to pick an exit status.
ErrUsage = errors.New("usage")
)
// CommandFlags is a function called before help or command execution to parse command line flags
// It should return a flag.FlagSet instance and a pointer to the struct that will contain parsed flags
type CommandFlags func() (*flag.FlagSet, any)
@@ -44,8 +55,10 @@ func execCommand(c *Command, args []string, w StringWriter) error {
fl.SetOutput(w.GetWriter())
err := fl.Parse(args)
if err != nil {
// fl.Parse has dumped error information to the user via the w writer.
return err
// fl.Parse has dumped error information to the user via the w writer, so
// the wrapper exists purely so a transport can tell a usage problem from a
// command that ran and failed.
return fmt.Errorf("%w: %w", ErrUsage, err)
}
args = fl.Args()
}
+226
View File
@@ -0,0 +1,226 @@
package diag
import (
"bufio"
"encoding/binary"
"encoding/json"
"errors"
"fmt"
"io"
)
// The ctl protocol is one request, one response, one connection.
//
// The request is a single JSON line. argv travels as a list rather than a joined string so
// that a path with a space in it survives the trip; the client already has a real argv from
// the operating system and re-splitting it would only ever lose information.
//
// The response is a stream of frames rather than raw bytes followed by a status line,
// because there is no sentinel that is safe to look for: `print-cert -raw` emits arbitrary
// PEM and `list-hostmap -json` emits arbitrary JSON, either of which could contain whatever
// terminator we picked.
const (
// ProtoVersion is the only request version this build understands. An unknown version
// gets a legible error rather than a hang, which is the whole point of sending it.
ProtoVersion = 1
// frameOutput carries raw command output, destined for the client's stdout.
frameOutput = 0x01
// frameEnd carries a JSON endPayload and is the last frame on a connection.
frameEnd = 0x02
// frameStderr is reserved. Commands write to a single writer today, so there is nothing
// to put in it, but holding the number means adding one later needs no version bump.
frameStderr = 0x03
// maxFrame bounds a single frame's payload. Larger writes are split across frames.
maxFrame = 64 * 1024
// maxRequest bounds the request line, so a client that never sends a newline cannot make
// nebula buffer without limit.
maxRequest = 64 * 1024
// outputBuffer is what keeps json.NewEncoder(w.GetWriter()) from emitting a frame per
// token; output accumulates here and flushes in useful sized chunks.
outputBuffer = 32 * 1024
)
// ErrTruncated means the connection ended before the end frame arrived, which is how a
// client notices that nebula died or was torn down partway through a command.
var ErrTruncated = errors.New("connection closed before the command finished")
// request is the JSON line a client sends.
type request struct {
Version int `json:"version"`
Args []string `json:"args"`
}
// endPayload is the JSON body of the end frame. Error is set only when Status is non-zero
// and describes a failure to run the command, not a failure the command itself reported.
type endPayload struct {
Status int `json:"status"`
Error string `json:"error,omitempty"`
}
// writeRequest sends the request line.
func writeRequest(w io.Writer, args []string) error {
b, err := json.Marshal(request{Version: ProtoVersion, Args: args})
if err != nil {
return err
}
if len(b)+1 > maxRequest {
return fmt.Errorf("command line is too long: %d bytes", len(b))
}
_, err = w.Write(append(b, '\n'))
return err
}
// readRequest reads and validates one request line.
func readRequest(r *bufio.Reader) (request, error) {
var req request
line, err := readLimitedLine(r, maxRequest)
if err != nil {
return req, err
}
if err := json.Unmarshal(line, &req); err != nil {
return req, fmt.Errorf("malformed request: %w", err)
}
if req.Version != ProtoVersion {
return req, fmt.Errorf("unsupported protocol version %d, this nebula speaks version %d", req.Version, ProtoVersion)
}
return req, nil
}
// readLimitedLine reads through the next newline, refusing a line longer than limit rather
// than buffering whatever an unfriendly client decides to send.
func readLimitedLine(r *bufio.Reader, limit int) ([]byte, error) {
line := make([]byte, 0, 256)
for {
b, err := r.ReadByte()
if err != nil {
return nil, err
}
if b == '\n' {
return line, nil
}
if len(line) >= limit {
return nil, fmt.Errorf("request exceeded %d bytes without a newline", limit)
}
line = append(line, b)
}
}
// frameWriter turns writes into output frames. It is handed to commands wrapped in a
// bufio.Writer, so a command that makes many small writes does not make many small frames.
type frameWriter struct {
w io.Writer
}
func (f *frameWriter) Write(b []byte) (int, error) {
written := 0
for {
chunk := b[written:]
if len(chunk) > maxFrame {
chunk = chunk[:maxFrame]
}
if err := writeFrame(f.w, frameOutput, chunk); err != nil {
return written, err
}
written += len(chunk)
if written == len(b) {
return written, nil
}
}
}
// writeFrame emits one frame: a type byte, a big endian length, then the payload.
func writeFrame(w io.Writer, kind byte, payload []byte) error {
var hdr [5]byte
hdr[0] = kind
binary.BigEndian.PutUint32(hdr[1:], uint32(len(payload)))
if _, err := w.Write(hdr[:]); err != nil {
return err
}
if len(payload) == 0 {
return nil
}
_, err := w.Write(payload)
return err
}
// writeEnd emits the final frame. A transport error here is unreportable by definition, the
// connection is the only channel we have.
func writeEnd(w io.Writer, status int, msg string) error {
b, err := json.Marshal(endPayload{Status: status, Error: msg})
if err != nil {
return err
}
return writeFrame(w, frameEnd, b)
}
// readResponse consumes frames until the end frame, copying output to out. It returns the
// command's exit status. A non-nil error means the exchange failed and the status is
// meaningless.
func readResponse(r io.Reader, out io.Writer) (int, error) {
var hdr [5]byte
for {
if _, err := io.ReadFull(r, hdr[:]); err != nil {
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) {
return 0, ErrTruncated
}
return 0, err
}
length := binary.BigEndian.Uint32(hdr[1:])
if length > maxFrame {
return 0, fmt.Errorf("frame of %d bytes exceeds the %d byte maximum", length, maxFrame)
}
payload := make([]byte, length)
if _, err := io.ReadFull(r, payload); err != nil {
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) {
return 0, ErrTruncated
}
return 0, err
}
switch hdr[0] {
case frameOutput:
if _, err := out.Write(payload); err != nil {
return 0, err
}
case frameEnd:
var end endPayload
if err := json.Unmarshal(payload, &end); err != nil {
return 0, fmt.Errorf("malformed end frame: %w", err)
}
if end.Error != "" {
return end.Status, errors.New(end.Error)
}
return end.Status, nil
case frameStderr:
// Reserved and unused by this build. Skipping rather than failing means an older
// client stays usable against a newer nebula that starts sending them.
default:
return 0, fmt.Errorf("unknown frame type 0x%02x", hdr[0])
}
}
}
+144
View File
@@ -0,0 +1,144 @@
package diag
import (
"bufio"
"bytes"
"encoding/binary"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestRequestRoundTrip(t *testing.T) {
t.Run("argv survives a round trip, spaces and all", func(t *testing.T) {
buf := &bytes.Buffer{}
args := []string{"start-cpu-profile", "/tmp/a path.pb.gz", "-json"}
require.NoError(t, writeRequest(buf, args))
req, err := readRequest(bufio.NewReader(buf))
require.NoError(t, err)
assert.Equal(t, ProtoVersion, req.Version)
assert.Equal(t, args, req.Args)
})
t.Run("an unknown version is refused by name", func(t *testing.T) {
r := bufio.NewReader(strings.NewReader(`{"version":99,"args":["version"]}` + "\n"))
_, err := readRequest(r)
require.Error(t, err)
assert.Contains(t, err.Error(), "unsupported protocol version 99")
})
t.Run("malformed json is refused", func(t *testing.T) {
r := bufio.NewReader(strings.NewReader("not json\n"))
_, err := readRequest(r)
require.Error(t, err)
assert.Contains(t, err.Error(), "malformed request")
})
t.Run("a line without a newline is bounded rather than buffered forever", func(t *testing.T) {
r := bufio.NewReader(strings.NewReader(strings.Repeat("a", maxRequest+10)))
_, err := readRequest(r)
require.Error(t, err)
assert.Contains(t, err.Error(), "without a newline")
})
}
func TestResponseRoundTrip(t *testing.T) {
t.Run("output and status survive a round trip", func(t *testing.T) {
wire := &bytes.Buffer{}
w := bufio.NewWriterSize(&frameWriter{w: wire}, outputBuffer)
require.NoError(t, NewWriter(w).WriteLine("hello"))
require.NoError(t, w.Flush())
require.NoError(t, writeEnd(wire, StatusOK, ""))
out := &bytes.Buffer{}
status, err := readResponse(wire, out)
require.NoError(t, err)
assert.Equal(t, StatusOK, status)
assert.Equal(t, "hello\n", out.String())
})
// print-cert -raw and list-hostmap -json both emit arbitrary bytes, so a payload larger
// than one frame has to reassemble exactly.
t.Run("a payload larger than one frame reassembles byte for byte", func(t *testing.T) {
big := bytes.Repeat([]byte("nebula"), maxFrame)
wire := &bytes.Buffer{}
fw := &frameWriter{w: wire}
n, err := fw.Write(big)
require.NoError(t, err)
require.Equal(t, len(big), n)
require.NoError(t, writeEnd(wire, StatusOK, ""))
out := &bytes.Buffer{}
status, err := readResponse(wire, out)
require.NoError(t, err)
assert.Equal(t, StatusOK, status)
assert.Equal(t, big, out.Bytes())
})
t.Run("a non-zero status carries its message", func(t *testing.T) {
wire := &bytes.Buffer{}
require.NoError(t, writeEnd(wire, StatusError, "it went wrong"))
status, err := readResponse(wire, &bytes.Buffer{})
require.Error(t, err)
assert.Equal(t, StatusError, status)
assert.Contains(t, err.Error(), "it went wrong")
})
// This is how the CLI notices a nebula that died mid-command rather than silently
// reporting whatever partial output it managed to read.
t.Run("a stream ending without an end frame is truncated, not successful", func(t *testing.T) {
wire := &bytes.Buffer{}
_, err := (&frameWriter{w: wire}).Write([]byte("partial"))
require.NoError(t, err)
out := &bytes.Buffer{}
_, err = readResponse(wire, out)
assert.ErrorIs(t, err, ErrTruncated)
})
t.Run("a truncated frame header is truncated, not successful", func(t *testing.T) {
_, err := readResponse(bytes.NewReader([]byte{frameOutput, 0x00}), &bytes.Buffer{})
assert.ErrorIs(t, err, ErrTruncated)
})
t.Run("an oversized frame is refused rather than allocated", func(t *testing.T) {
var hdr [5]byte
hdr[0] = frameOutput
binary.BigEndian.PutUint32(hdr[1:], maxFrame+1)
_, err := readResponse(bytes.NewReader(hdr[:]), &bytes.Buffer{})
require.Error(t, err)
assert.Contains(t, err.Error(), "exceeds")
})
// A reserved frame an older client does not understand must not break it.
t.Run("a reserved frame type is skipped", func(t *testing.T) {
wire := &bytes.Buffer{}
require.NoError(t, writeFrame(wire, frameStderr, []byte("future")))
require.NoError(t, writeFrame(wire, frameOutput, []byte("now")))
require.NoError(t, writeEnd(wire, StatusOK, ""))
out := &bytes.Buffer{}
status, err := readResponse(wire, out)
require.NoError(t, err)
assert.Equal(t, StatusOK, status)
assert.Equal(t, "now", out.String())
})
t.Run("an unknown frame type is an error", func(t *testing.T) {
wire := &bytes.Buffer{}
require.NoError(t, writeFrame(wire, 0x7f, nil))
_, err := readResponse(wire, &bytes.Buffer{})
require.Error(t, err)
assert.Contains(t, err.Error(), "unknown frame type")
})
}
+125
View File
@@ -0,0 +1,125 @@
package diag
import (
"fmt"
"sync"
"github.com/anmitsu/go-shlex"
"github.com/armon/go-radix"
)
// Registry is the set of commands nebula exposes for debugging and administration. It is
// transport neutral: the ssh console and the `nebula ctl` unix socket dispatch against the
// same registry, and neither knows the other exists.
//
// Registration is expected to happen once during startup, before any transport is serving,
// but the lock makes a late RegisterCommand safe rather than a data race waiting to happen.
type Registry struct {
mu sync.RWMutex
commands *radix.Tree
}
// NewRegistry returns a registry containing only `help`. Everything else is attached by
// the caller, see attachCommands in the nebula package.
func NewRegistry() *Registry {
r := &Registry{commands: radix.New()}
r.RegisterCommand(&Command{
Name: "help",
ShortDescription: "prints available commands or help <command> for specific usage info",
Callback: func(a any, args []string, w StringWriter) error {
return r.help(args, w)
},
})
return r
}
// RegisterCommand adds a command that a user can run.
func (r *Registry) RegisterCommand(c *Command) {
r.mu.Lock()
defer r.mu.Unlock()
r.commands.Insert(c.Name, c)
}
// Clone returns an independent copy sharing no tree with the original. The ssh session uses
// this so the `logout` command it adds for itself is invisible to every other session, and
// to `nebula ctl`.
func (r *Registry) Clone() *Registry {
r.mu.RLock()
defer r.mu.RUnlock()
return &Registry{commands: radix.NewFromMap(r.commands.ToMap())}
}
// Match returns every registered command name carrying the given prefix, for tab completion.
func (r *Registry) Match(prefix string) []string {
r.mu.RLock()
defer r.mu.RUnlock()
return matchCommand(r.commands, prefix)
}
// Dispatch splits line the way a shell would and runs the result. The ssh console uses this
// because a terminal only ever hands it a line; a transport that already has a real argv
// should call DispatchArgs instead rather than round tripping through a quoting parser.
func (r *Registry) Dispatch(line string, w StringWriter) error {
args, err := shlex.Split(line, true)
if err != nil {
if wErr := w.WriteLine(fmt.Sprintf("Unable to parse command: %s", err)); wErr != nil {
return wErr
}
return err
}
return r.DispatchArgs(args, w)
}
// DispatchArgs runs args[0] with args[1:] as its arguments, writing everything the command
// produces to w. An empty args dumps the command list, matching what an empty line does on
// the ssh console.
//
// Callbacks report user facing problems as prose on w and return nil by convention, so a
// non-nil error here means the command could not be run at all: ErrUnknownCommand, an
// ErrUsage wrapped flag failure, or an internal failure a callback chose to surface.
func (r *Registry) DispatchArgs(args []string, w StringWriter) error {
if len(args) == 0 {
r.mu.RLock()
defer r.mu.RUnlock()
dumpCommands(r.commands, w)
return nil
}
r.mu.RLock()
cmd, err := lookupCommand(r.commands, args[0])
r.mu.RUnlock()
if err != nil {
if wErr := w.WriteLine(fmt.Sprintf("Command lookup failed: %s", err)); wErr != nil {
return wErr
}
return err
}
if cmd == nil {
if wErr := w.WriteLine(fmt.Sprintf("Did not understand: %s", args[0])); wErr != nil {
return wErr
}
r.mu.RLock()
defer r.mu.RUnlock()
dumpCommands(r.commands, w)
return fmt.Errorf("%w: %s", ErrUnknownCommand, args[0])
}
// -h and -help anywhere in the arguments mean the user wants to know how the command
// works, not to run it.
if checkHelpArgs(args) {
return r.help([]string{cmd.Name}, w)
}
return execCommand(cmd, args[1:], w)
}
// help renders the command list, or one command's usage, onto w.
func (r *Registry) help(args []string, w StringWriter) error {
r.mu.RLock()
defer r.mu.RUnlock()
return helpCallback(r.commands, args, w)
}
+167
View File
@@ -0,0 +1,167 @@
package diag
import (
"bytes"
"flag"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
type testFlags struct {
Json bool
}
// testCommand builds a command carrying a flag set, recording what the callback was actually
// handed so a test can assert on it.
func testCommand(name string, seen *any, args *[]string) *Command {
return &Command{
Name: name,
ShortDescription: name + " short description",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
f := &testFlags{}
fl.BoolVar(&f.Json, "json", false, "outputs json")
return fl, f
},
Callback: func(fs any, a []string, w StringWriter) error {
if seen != nil {
*seen = fs
}
if args != nil {
*args = a
}
return w.WriteLine("ran " + name)
},
}
}
func newTestRegistry(t *testing.T) (*Registry, *bytes.Buffer, StringWriter) {
t.Helper()
buf := &bytes.Buffer{}
return NewRegistry(), buf, NewWriter(buf)
}
func TestRegistryDispatch(t *testing.T) {
t.Run("a new registry knows help and nothing else", func(t *testing.T) {
r, buf, w := newTestRegistry(t)
require.NoError(t, r.DispatchArgs([]string{"help"}, w))
assert.Contains(t, buf.String(), "help -")
})
t.Run("empty args dump the command list, matching an empty line on the console", func(t *testing.T) {
r, buf, w := newTestRegistry(t)
r.RegisterCommand(testCommand("do-thing", nil, nil))
require.NoError(t, r.DispatchArgs(nil, w))
assert.Contains(t, buf.String(), "Available commands:")
assert.Contains(t, buf.String(), "do-thing - do-thing short description")
})
t.Run("an unknown command reports ErrUnknownCommand and still tells the user", func(t *testing.T) {
r, buf, w := newTestRegistry(t)
err := r.DispatchArgs([]string{"nope"}, w)
require.ErrorIs(t, err, ErrUnknownCommand)
assert.Contains(t, buf.String(), "Did not understand: nope")
assert.Contains(t, buf.String(), "Available commands:")
})
// This is the hazard the ctl transport has to preserve: every callback in ssh.go begins by
// type asserting fs to its own concrete flags struct. Reach a callback without going
// through Command.Flags and every one of them fails.
t.Run("a callback is handed the concrete struct its Flags callback returned", func(t *testing.T) {
var seen any
r, _, w := newTestRegistry(t)
r.RegisterCommand(testCommand("do-thing", &seen, nil))
require.NoError(t, r.DispatchArgs([]string{"do-thing", "-json"}, w))
flags, ok := seen.(*testFlags)
require.True(t, ok, "callback was handed %T, not *testFlags", seen)
assert.True(t, flags.Json)
})
t.Run("positional arguments survive flag parsing", func(t *testing.T) {
var args []string
r, _, w := newTestRegistry(t)
r.RegisterCommand(testCommand("do-thing", nil, &args))
require.NoError(t, r.DispatchArgs([]string{"do-thing", "-json", "10.0.0.1"}, w))
assert.Equal(t, []string{"10.0.0.1"}, args)
})
// Documents stdlib flag behaviour rather than endorsing it: parsing stops at the first
// positional, so a flag written after one is silently a positional too.
t.Run("a flag after a positional is not parsed as a flag", func(t *testing.T) {
var seen any
var args []string
r, _, w := newTestRegistry(t)
r.RegisterCommand(testCommand("do-thing", &seen, &args))
require.NoError(t, r.DispatchArgs([]string{"do-thing", "10.0.0.1", "-json"}, w))
assert.False(t, seen.(*testFlags).Json)
assert.Equal(t, []string{"10.0.0.1", "-json"}, args)
})
t.Run("a bad flag reports ErrUsage and writes the usage text", func(t *testing.T) {
r, buf, w := newTestRegistry(t)
r.RegisterCommand(testCommand("do-thing", nil, nil))
err := r.DispatchArgs([]string{"do-thing", "-nope"}, w)
require.ErrorIs(t, err, ErrUsage)
assert.Contains(t, buf.String(), "flag provided but not defined")
})
t.Run("-h anywhere routes to help instead of running the command", func(t *testing.T) {
var seen any
r, buf, w := newTestRegistry(t)
r.RegisterCommand(testCommand("do-thing", &seen, nil))
require.NoError(t, r.DispatchArgs([]string{"do-thing", "-h"}, w))
assert.Nil(t, seen, "the callback should not have run")
assert.Contains(t, buf.String(), "do-thing - do-thing short description")
assert.Contains(t, buf.String(), "-json")
})
t.Run("Dispatch splits a line the way a shell would", func(t *testing.T) {
var args []string
r, _, w := newTestRegistry(t)
r.RegisterCommand(testCommand("do-thing", nil, &args))
require.NoError(t, r.Dispatch(`do-thing "/tmp/a path.pb.gz"`, w))
assert.Equal(t, []string{"/tmp/a path.pb.gz"}, args)
})
t.Run("Match returns names by prefix for tab completion", func(t *testing.T) {
r, _, _ := newTestRegistry(t)
r.RegisterCommand(testCommand("print-cert", nil, nil))
r.RegisterCommand(testCommand("print-tunnel", nil, nil))
r.RegisterCommand(testCommand("version", nil, nil))
assert.Equal(t, []string{"print-cert", "print-tunnel"}, r.Match("print-"))
})
}
// A clone is what keeps the ssh session's `logout` command from being visible to every other
// session, and to nebula ctl.
func TestRegistryCloneIsolation(t *testing.T) {
parent, _, w := newTestRegistry(t)
parent.RegisterCommand(testCommand("shared", nil, nil))
child := parent.Clone()
child.RegisterCommand(testCommand("logout", nil, nil))
require.NoError(t, child.DispatchArgs([]string{"logout"}, w))
buf := &bytes.Buffer{}
err := parent.DispatchArgs([]string{"logout"}, NewWriter(buf))
assert.ErrorIs(t, err, ErrUnknownCommand)
buf.Reset()
require.NoError(t, child.DispatchArgs([]string{"shared"}, NewWriter(buf)))
assert.True(t, strings.HasPrefix(buf.String(), "ran shared"))
}
+137
View File
@@ -0,0 +1,137 @@
package diag
import (
"bufio"
"context"
"errors"
"fmt"
"log/slog"
"net"
"time"
)
// Exit statuses the client reports. They follow shell convention closely enough that a
// script can tell "you asked for something that does not exist" from "it ran and failed".
const (
// StatusOK means the command ran. Note that commands report their own user facing
// problems as prose and still exit 0, matching the ssh console.
StatusOK = 0
// StatusError means the command could not be completed.
StatusError = 1
// StatusUsage means the arguments were not valid for that command.
StatusUsage = 2
// StatusUnknownCommand means there is no such command.
StatusUnknownCommand = 127
)
// requestTimeout bounds how long a connected client may take to send its request line. There
// is deliberately no timeout on the response: `reload` runs every reload callback inline
// before it returns, and a slow one is not a reason to hang up on the operator.
const requestTimeout = 5 * time.Second
// Server serves a Registry over a stream listener. It knows nothing about unix sockets, so
// tests can drive it over a net.Pipe.
type Server struct {
l *slog.Logger
reg *Registry
}
func NewServer(l *slog.Logger, reg *Registry) *Server {
return &Server{l: l, reg: reg}
}
// Serve accepts connections until ln is closed. Cancelling ctx closes ln, which is what ends
// the accept loop; a listener closed underneath us is a normal shutdown, not an error.
func (s *Server) Serve(ctx context.Context, ln net.Listener) error {
go func() {
<-ctx.Done()
if err := ln.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
s.l.Warn("Failed to close the ctl listener", "error", err)
}
}()
for {
conn, err := ln.Accept()
if err != nil {
if errors.Is(err, net.ErrClosed) || ctx.Err() != nil {
return nil
}
return err
}
go s.ServeConn(ctx, conn)
}
}
// ServeConn handles one request and closes c.
func (s *Server) ServeConn(ctx context.Context, c net.Conn) {
defer func() {
if err := c.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
s.l.Debug("Failed to close a ctl connection", "error", err)
}
}()
if err := c.SetReadDeadline(time.Now().Add(requestTimeout)); err != nil {
s.l.Debug("Failed to set a ctl read deadline", "error", err)
}
req, err := readRequest(bufio.NewReaderSize(c, maxRequest))
if err != nil {
s.l.Debug("Rejected a ctl request", "error", err)
// Best effort: the client may already be gone, and there is nowhere else to report it.
_ = writeEnd(c, StatusError, err.Error())
return
}
// The request is in hand, so the command owns the rest of the connection's lifetime.
if err := c.SetReadDeadline(time.Time{}); err != nil {
s.l.Debug("Failed to clear the ctl read deadline", "error", err)
}
s.l.Debug("Running a ctl command", "args", req.Args)
buf := bufio.NewWriterSize(&frameWriter{w: c}, outputBuffer)
dispatchErr := s.reg.DispatchArgs(req.Args, NewWriter(buf))
if err := buf.Flush(); err != nil {
s.l.Debug("Failed to flush ctl output", "error", err)
return
}
status, msg := statusFor(dispatchErr)
if err := writeEnd(c, status, msg); err != nil {
s.l.Debug("Failed to write the ctl end frame", "error", err)
}
}
// StatusFor maps a dispatch error onto an exit status, for a transport that has somewhere to
// put one.
func StatusFor(err error) int {
status, _ := statusFor(err)
return status
}
// statusFor maps a dispatch error onto an exit status and, when the failure is ours to
// explain rather than one the command already wrote as prose, a message to go with it.
func statusFor(err error) (int, string) {
switch {
case err == nil:
return StatusOK, ""
case errors.Is(err, ErrUnknownCommand):
return StatusUnknownCommand, ""
case errors.Is(err, ErrUsage):
return StatusUsage, ""
default:
return StatusError, fmt.Sprintf("%s", err)
}
}
// ErrNotSupported means this platform has no ctl transport. Windows is waiting on a named
// pipe implementation; mobile has no daemon for a CLI to attach to in the first place.
var ErrNotSupported = errors.New("nebula ctl is not supported on this platform")
// Listen creates the ctl listener at path. It is the platform boundary: everything above it
// in this package is portable.
func Listen(path string) (net.Listener, error) {
return listenSocket(path)
}
+273
View File
@@ -0,0 +1,273 @@
//go:build !windows
package diag
import (
"bytes"
"context"
"errors"
"fmt"
"io/fs"
"log/slog"
"net"
"os"
"path/filepath"
"sync"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// testSocketPath returns a short socket path. t.TempDir on darwin lives under
// /var/folders/... and readily exceeds the 104 byte sun_path limit, which fails as a bare
// "invalid argument" a long way from the cause.
func testSocketPath(t *testing.T) string {
t.Helper()
dir, err := os.MkdirTemp("/tmp", "nebctl")
require.NoError(t, err)
t.Cleanup(func() { _ = os.RemoveAll(dir) })
path := filepath.Join(dir, "sub", "ctl.sock")
require.LessOrEqual(t, len(path), maxSocketPath, "test socket path is too long for sun_path")
return path
}
func newTestServer(t *testing.T) (*Registry, string) {
t.Helper()
reg := NewRegistry()
path := testSocketPath(t)
ln, err := Listen(path)
require.NoError(t, err)
ctx, cancel := context.WithCancel(context.Background())
srv := NewServer(slog.New(slog.DiscardHandler), reg)
var wg sync.WaitGroup
wg.Add(1)
go func() {
defer wg.Done()
assert.NoError(t, srv.Serve(ctx, ln))
}()
t.Cleanup(func() {
cancel()
wg.Wait()
})
return reg, path
}
func run(t *testing.T, path string, args ...string) (string, int, error) {
t.Helper()
c, err := Dial(path)
require.NoError(t, err)
defer c.Close()
out := &bytes.Buffer{}
status, err := c.Run(args, out)
return out.String(), status, err
}
func TestServeConn(t *testing.T) {
t.Run("a command runs and its output comes back", func(t *testing.T) {
reg, path := newTestServer(t)
reg.RegisterCommand(&Command{
Name: "version",
ShortDescription: "prints a version",
Callback: func(fs any, a []string, w StringWriter) error {
return w.WriteLine("1.2.3")
},
})
out, status, err := run(t, path, "version")
require.NoError(t, err)
assert.Equal(t, StatusOK, status)
assert.Equal(t, "1.2.3\n", out)
})
t.Run("no args gets the command list", func(t *testing.T) {
_, path := newTestServer(t)
out, status, err := run(t, path)
require.NoError(t, err)
assert.Equal(t, StatusOK, status)
assert.Contains(t, out, "Available commands:")
})
t.Run("an unknown command exits 127", func(t *testing.T) {
_, path := newTestServer(t)
out, status, err := run(t, path, "nope")
require.NoError(t, err)
assert.Equal(t, StatusUnknownCommand, status)
assert.Contains(t, out, "Did not understand: nope")
})
t.Run("a bad flag exits 2", func(t *testing.T) {
reg, path := newTestServer(t)
var seen any
reg.RegisterCommand(testCommand("do-thing", &seen, nil))
out, status, err := run(t, path, "do-thing", "-nope")
require.NoError(t, err)
assert.Equal(t, StatusUsage, status)
assert.Contains(t, out, "flag provided but not defined")
})
t.Run("a callback error exits 1 and reports why", func(t *testing.T) {
reg, path := newTestServer(t)
reg.RegisterCommand(&Command{
Name: "explode",
ShortDescription: "fails",
Callback: func(fs any, a []string, w StringWriter) error {
return errors.New("boom")
},
})
_, status, err := run(t, path, "explode")
require.Error(t, err)
assert.Equal(t, StatusError, status)
assert.Contains(t, err.Error(), "boom")
})
t.Run("output larger than the buffer arrives intact", func(t *testing.T) {
reg, path := newTestServer(t)
want := bytes.Repeat([]byte("x"), outputBuffer*3+7)
reg.RegisterCommand(&Command{
Name: "big",
ShortDescription: "writes a lot",
Callback: func(fs any, a []string, w StringWriter) error {
return w.WriteBytes(want)
},
})
out, status, err := run(t, path, "big")
require.NoError(t, err)
assert.Equal(t, StatusOK, status)
assert.Equal(t, string(want), out)
})
t.Run("concurrent clients are all served", func(t *testing.T) {
reg, path := newTestServer(t)
reg.RegisterCommand(&Command{
Name: "slow",
ShortDescription: "takes a moment",
Callback: func(fs any, a []string, w StringWriter) error {
time.Sleep(10 * time.Millisecond)
return w.WriteLine("done")
},
})
var wg sync.WaitGroup
for i := 0; i < 8; i++ {
wg.Add(1)
go func() {
defer wg.Done()
out, status, err := run(t, path, "slow")
assert.NoError(t, err)
assert.Equal(t, StatusOK, status)
assert.Equal(t, "done\n", out)
}()
}
wg.Wait()
})
t.Run("a client that hangs up mid command does not take the server down", func(t *testing.T) {
reg, path := newTestServer(t)
reg.RegisterCommand(&Command{
Name: "version",
ShortDescription: "prints a version",
Callback: func(fs any, a []string, w StringWriter) error {
return w.WriteLine("1.2.3")
},
})
c, err := Dial(path)
require.NoError(t, err)
require.NoError(t, writeRequest(c.conn, []string{"version"}))
require.NoError(t, c.Close())
// The next client still gets served.
out, status, err := run(t, path, "version")
require.NoError(t, err)
assert.Equal(t, StatusOK, status)
assert.Equal(t, "1.2.3\n", out)
})
}
func TestListenSocket(t *testing.T) {
t.Run("the socket is 0600 inside a 0700 directory", func(t *testing.T) {
path := testSocketPath(t)
ln, err := Listen(path)
require.NoError(t, err)
defer ln.Close()
fi, err := os.Stat(path)
require.NoError(t, err)
assert.Equal(t, os.FileMode(0600), fi.Mode().Perm(), "socket mode")
di, err := os.Stat(filepath.Dir(path))
require.NoError(t, err)
assert.Equal(t, os.FileMode(0700), di.Mode().Perm(), "socket directory mode")
})
t.Run("the socket is unlinked when the listener closes", func(t *testing.T) {
path := testSocketPath(t)
ln, err := Listen(path)
require.NoError(t, err)
require.NoError(t, ln.Close())
_, err = os.Stat(path)
assert.ErrorIs(t, err, fs.ErrNotExist)
})
// A crashed nebula leaves its socket behind, and the next one has to be able to start.
t.Run("a socket left behind by a dead nebula is replaced", func(t *testing.T) {
path := testSocketPath(t)
ln, err := Listen(path)
require.NoError(t, err)
// Close the listener without unlinking, the way a killed process leaves things.
unix, ok := ln.(*net.UnixListener)
require.True(t, ok)
unix.SetUnlinkOnClose(false)
require.NoError(t, ln.Close())
require.FileExists(t, path)
ln2, err := Listen(path)
require.NoError(t, err)
assert.NoError(t, ln2.Close())
})
// Silently stealing it would break the nebula that got there first.
t.Run("a socket another nebula is serving is refused", func(t *testing.T) {
_, path := newTestServer(t)
_, err := Listen(path)
require.Error(t, err)
assert.Contains(t, err.Error(), "already being served")
})
t.Run("a path that is not a socket is refused rather than removed", func(t *testing.T) {
path := testSocketPath(t)
require.NoError(t, os.MkdirAll(filepath.Dir(path), 0700))
require.NoError(t, os.WriteFile(path, []byte("precious"), 0600))
_, err := Listen(path)
require.Error(t, err)
assert.Contains(t, err.Error(), "is not a socket")
assert.FileExists(t, path, "the file must not have been removed")
})
t.Run("a path too long for sun_path says so", func(t *testing.T) {
_, err := Listen("/tmp/" + fmt.Sprintf("%0*d", maxSocketPath, 0) + "/ctl.sock")
require.Error(t, err)
assert.Contains(t, err.Error(), "the maximum is")
})
}
+107
View File
@@ -0,0 +1,107 @@
//go:build !windows
package diag
import (
"errors"
"fmt"
"io/fs"
"net"
"os"
"path/filepath"
"runtime"
"time"
)
// maxSocketPath is the smallest sun_path across the platforms nebula ships on: 104 bytes on
// darwin and the BSDs, 108 on Linux. Checking it ourselves turns a bare "invalid argument"
// into something an operator can act on.
const maxSocketPath = 103
// DefaultSocketPath is where nebula listens when ctl.socket is unset. An empty string means
// the platform has no sensible default and ctl stays off unless an operator names a path.
func DefaultSocketPath() string {
switch runtime.GOOS {
case "ios", "android":
// No daemon to attach to and no shell to attach from, and nowhere writable that
// would survive being guessed. Mobile embedders drive nebula through Control.
return ""
case "linux":
return "/run/nebula/ctl.sock"
default:
// /run does not exist on darwin, and /var/run is the portable spelling everywhere
// else nebula builds.
return "/var/run/nebula/ctl.sock"
}
}
// listenSocket creates the listening socket at path, taking over one a previous nebula left
// behind but refusing one that is still being served.
func listenSocket(path string) (net.Listener, error) {
if len(path) > maxSocketPath {
return nil, fmt.Errorf("socket path is %d bytes, the maximum is %d", len(path), maxSocketPath)
}
// The directory, not the socket, is what enforces access control. net.Listen creates the
// socket with 0777&^umask, so with a typical 0022 umask it is world connectable for the
// window between bind and chmod. Nobody can traverse into a 0700 directory to reach it in
// that window, and unlike the socket's own mode, directory traversal is enforced
// consistently across every platform this file builds for.
dir := filepath.Dir(path)
if err := os.MkdirAll(dir, 0700); err != nil {
return nil, fmt.Errorf("failed to create %s: %w", dir, err)
}
if err := os.Chmod(dir, 0700); err != nil {
return nil, fmt.Errorf("failed to set permissions on %s: %w", dir, err)
}
if err := clearStaleSocket(path); err != nil {
return nil, err
}
ln, err := net.Listen("unix", path)
if err != nil {
return nil, err
}
// Defence in depth behind the directory, for anyone who relocates the socket somewhere
// more permissive.
if err := os.Chmod(path, 0600); err != nil {
_ = ln.Close()
return nil, fmt.Errorf("failed to set permissions on %s: %w", path, err)
}
return ln, nil
}
// dialSocket connects to a nebula serving at path.
func dialSocket(path string, timeout time.Duration) (net.Conn, error) {
return net.DialTimeout("unix", path, timeout)
}
// clearStaleSocket removes a socket a crashed nebula left behind, but refuses to steal one
// another nebula is still serving. Two instances on one host need two paths; they cannot
// share one, and silently taking the socket would break the instance that got there first.
func clearStaleSocket(path string) error {
fi, err := os.Lstat(path)
if errors.Is(err, fs.ErrNotExist) {
return nil
}
if err != nil {
return err
}
if fi.Mode()&fs.ModeSocket == 0 {
return fmt.Errorf("%s exists and is not a socket, refusing to remove it", path)
}
// A successful dial is the only reliable way to tell a live socket from an abandoned
// one; the inode looks identical either way.
c, err := net.DialTimeout("unix", path, 100*time.Millisecond)
if err == nil {
_ = c.Close()
return fmt.Errorf("%s is already being served, is another nebula running?", path)
}
return os.Remove(path)
}
+27
View File
@@ -0,0 +1,27 @@
//go:build windows
package diag
import (
"net"
"time"
)
// Windows has AF_UNIX since Windows 10 1803, but no way to secure the socket that resembles
// what the unix build does: os.Chmod cannot express an ACL, and a socket's reachability comes
// down to whatever its directory inherited. Doing this properly means a named pipe with an
// explicit security descriptor, which is a dependency and a design this change does not carry.
// Until then the stub keeps the package building and gives operators a real answer.
// DefaultSocketPath returns an empty string: there is no path worth defaulting to here.
func DefaultSocketPath() string {
return ""
}
func listenSocket(path string) (net.Listener, error) {
return nil, ErrNotSupported
}
func dialSocket(path string, timeout time.Duration) (net.Conn, error) {
return nil, ErrNotSupported
}
+7 -1
View File
@@ -1,4 +1,4 @@
package sshd
package diag
import "io"
@@ -30,3 +30,9 @@ func (w *stringWriter) WriteBytes(b []byte) error {
func (w *stringWriter) GetWriter() io.Writer {
return w.w
}
// NewWriter adapts an io.Writer to the StringWriter commands are handed. Transports
// implement their own framing behind w; the commands never know the difference.
func NewWriter(w io.Writer) StringWriter {
return &stringWriter{w: w}
}
+6
View File
@@ -116,6 +116,9 @@ func newSimpleServerWithUdpAndUnsafeNetworks(v cert.Version, caCrt cert.Certific
"key": string(myPrivKey),
},
//"tun": m{"disabled": true},
// Several tests bring up more than one nebula in this process, and they would all
// contend for the same default ctl socket path. None of them exercise it.
"ctl": m{"enabled": false},
"firewall": m{
"outbound": []m{{
"proto": "any",
@@ -213,6 +216,9 @@ func newServer(caCrt []cert.Certificate, certs []cert.Certificate, key []byte, o
"key": string(key),
},
//"tun": m{"disabled": true},
// Several tests bring up more than one nebula in this process, and they would all
// contend for the same default ctl socket path. None of them exercise it.
"ctl": m{"enabled": false},
"firewall": m{
"outbound": []m{{
"proto": "any",
+57
View File
@@ -223,3 +223,60 @@ 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.
// Long connection manager timers so it never fires a direct test packet at the relay tunnel and bumps its
// epoch mid-test, which is the only other thing that touches that tunnel and would flake the assertion below.
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.1/24",
m{"relay": m{"use_relays": true}, "timers": m{"connection_alive_interval": 3600, "pending_deletion_interval": 3600}})
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()
}
+24
View File
@@ -236,6 +236,30 @@ punchy:
# Overriding this to "" is the same as "/" and will allow overwriting any path on the host.
#sandbox_dir: /var/tmp/nebula-debug
# ctl exposes nebula's debug and administrative commands over a local unix socket, so that `nebula ctl <command>` can
# reach the same commands the sshd block offers above without running an ssh server. Run `nebula ctl` on its own for the
# list of commands. Anyone who can open the socket can do everything the ssh console can, including closing tunnels,
# changing remotes, and writing profile data to disk, so the socket lives in a directory only the user nebula runs as
# can enter. Enabled by default. Not supported on Windows yet, and never enabled on iOS or Android.
#ctl:
# Toggles the feature. This setting is reloadable.
#enabled: true
# socket is the unix socket to listen on. The parent directory is created if it is missing and made readable only by
# the user nebula runs as, and a socket left behind by a crashed nebula is replaced. Defaults to /run/nebula/ctl.sock
# on Linux and /var/run/nebula/ctl.sock everywhere else; running nebula as a non-root user means picking a path it can
# write. Two nebulas on one host need two paths, the second to start will log that the socket is already being served
# and carry on without one. `nebula ctl` reads this value from the same config file when it is given -config, and
# otherwise assumes the default above. This setting is reloadable.
#socket: /run/nebula/ctl.sock
# sandbox_dir restricts the file paths the profiling commands (start-cpu-profile, save-heap-profile,
# save-mutex-profile) may write, exactly like sshd.sandbox_dir above, which it defaults to. Note that these paths are
# resolved by the nebula process and not by the shell running `nebula ctl`, so a relative path lands in this directory
# rather than in your working directory, and under a systemd unit with PrivateTmp=yes it lands somewhere your shell
# cannot see at all. The directory is NOT automatically created.
#sandbox_dir: /var/tmp/nebula-debug
# 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.
+15 -14
View File
@@ -21,6 +21,7 @@ import (
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/iputil"
)
type FirewallInterface interface {
@@ -262,11 +263,11 @@ func (f *Firewall) AddRule(incoming bool, proto uint8, startPort int32, endPort
}
switch proto {
case firewall.ProtoTCP:
case iputil.IPProtocolTCP:
fp = ft.TCP
case firewall.ProtoUDP:
case iputil.IPProtocolUDP:
fp = ft.UDP
case firewall.ProtoICMP, firewall.ProtoICMPv6:
case iputil.IPProtocolICMP, iputil.IPProtocolICMPv6:
//ICMP traffic doesn't have ports, so we always coerce to "any", even if a value is provided
if startPort != firewall.PortAny {
f.l.Warn("ignoring port specification for ICMP firewall rule", "startPort", startPort)
@@ -364,13 +365,13 @@ func AddFirewallRulesFromConfig(l *slog.Logger, inbound bool, c *config.C, fw Fi
proto = firewall.ProtoAny
startPort, endPort, err = parsePort(sPort)
case "tcp":
proto = firewall.ProtoTCP
proto = iputil.IPProtocolTCP
startPort, endPort, err = parsePort(sPort)
case "udp":
proto = firewall.ProtoUDP
proto = iputil.IPProtocolUDP
startPort, endPort, err = parsePort(sPort)
case "icmp":
proto = firewall.ProtoICMP
proto = iputil.IPProtocolICMP
startPort = firewall.PortAny
endPort = firewall.PortAny
if sPort != "" {
@@ -560,9 +561,9 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
}
switch fp.Protocol {
case firewall.ProtoTCP:
case iputil.IPProtocolTCP:
c.Expires = time.Now().Add(f.TCPTimeout)
case firewall.ProtoUDP:
case iputil.IPProtocolUDP:
c.Expires = time.Now().Add(f.UDPTimeout)
default:
c.Expires = time.Now().Add(f.DefaultTimeout)
@@ -582,9 +583,9 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
c := &conn{}
switch fp.Protocol {
case firewall.ProtoTCP:
case iputil.IPProtocolTCP:
timeout = f.TCPTimeout
case firewall.ProtoUDP:
case iputil.IPProtocolUDP:
timeout = f.UDPTimeout
default:
timeout = f.DefaultTimeout
@@ -635,15 +636,15 @@ func (ft *FirewallTable) match(p firewall.Packet, incoming bool, c *cert.CachedC
}
switch p.Protocol {
case firewall.ProtoTCP:
case iputil.IPProtocolTCP:
if ft.TCP.match(p, incoming, c, caPool) {
return true
}
case firewall.ProtoUDP:
case iputil.IPProtocolUDP:
if ft.UDP.match(p, incoming, c, caPool) {
return true
}
case firewall.ProtoICMP, firewall.ProtoICMPv6:
case iputil.IPProtocolICMP, iputil.IPProtocolICMPv6:
if ft.ICMP.match(p, incoming, c, caPool) {
return true
}
@@ -680,7 +681,7 @@ func (fp firewallPort) match(p firewall.Packet, incoming bool, c *cert.CachedCer
}
// this branch is here to catch traffic from FirewallTable.Any.match and FirewallTable.ICMP.match
if p.Protocol == firewall.ProtoICMP || p.Protocol == firewall.ProtoICMPv6 {
if p.Protocol == iputil.IPProtocolICMP || p.Protocol == iputil.IPProtocolICMPv6 {
// port numbers are re-used for connection tracking of ICMP,
// but we don't want to actually filter on them.
return fp[firewall.PortAny].match(p, c, caPool)
+7 -10
View File
@@ -4,17 +4,14 @@ import (
"encoding/json"
"fmt"
"net/netip"
"github.com/slackhq/nebula/iputil"
)
type m = map[string]any
const (
ProtoAny = 0 // When we want to handle HOPOPT (0) we can change this, if ever
ProtoTCP = 6
ProtoUDP = 17
ProtoICMP = 1
ProtoICMPv6 = 58
ProtoAny = 0 // When we want to handle HOPOPT (0) we can change this, if ever
PortAny = 0 // Special value for matching `port: any`
PortFragment = -1 // Special value for matching `port: fragment`
)
@@ -45,13 +42,13 @@ func (fp *Packet) Copy() *Packet {
func (fp Packet) MarshalJSON() ([]byte, error) {
var proto string
switch fp.Protocol {
case ProtoTCP:
case iputil.IPProtocolTCP:
proto = "tcp"
case ProtoICMP:
case iputil.IPProtocolICMP:
proto = "icmp"
case ProtoICMPv6:
case iputil.IPProtocolICMPv6:
proto = "icmpv6"
case ProtoUDP:
case iputil.IPProtocolUDP:
proto = "udp"
default:
proto = fmt.Sprintf("unknown %v", fp.Protocol)
+34 -33
View File
@@ -13,6 +13,7 @@ import (
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/iputil"
"github.com/slackhq/nebula/test"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -72,20 +73,20 @@ func TestFirewall_AddRule(t *testing.T) {
ti6, err := netip.ParsePrefix("fd12::34/128")
require.NoError(t, err)
require.NoError(t, fw.AddRule(true, firewall.ProtoTCP, 1, 1, []string{}, "", "", "", "", ""))
require.NoError(t, fw.AddRule(true, iputil.IPProtocolTCP, 1, 1, []string{}, "", "", "", "", ""))
// An empty rule is any
assert.True(t, fw.InRules.TCP[1].Any.Any.Any)
assert.Empty(t, fw.InRules.TCP[1].Any.Groups)
assert.Empty(t, fw.InRules.TCP[1].Any.Hosts)
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
require.NoError(t, fw.AddRule(true, firewall.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "", ""))
require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "", ""))
assert.Nil(t, fw.InRules.UDP[1].Any.Any)
assert.Contains(t, fw.InRules.UDP[1].Any.Groups[0].Groups, "g1")
assert.Empty(t, fw.InRules.UDP[1].Any.Hosts)
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
require.NoError(t, fw.AddRule(true, firewall.ProtoICMP, 1, 1, []string{}, "h1", "", "", "", ""))
require.NoError(t, fw.AddRule(true, iputil.IPProtocolICMP, 1, 1, []string{}, "h1", "", "", "", ""))
//no matter what port is given for icmp, it should end up as "any"
assert.Nil(t, fw.InRules.ICMP[firewall.PortAny].Any.Any)
assert.Empty(t, fw.InRules.ICMP[firewall.PortAny].Any.Groups)
@@ -116,11 +117,11 @@ func TestFirewall_AddRule(t *testing.T) {
assert.True(t, ok)
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
require.NoError(t, fw.AddRule(true, firewall.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "ca-name", ""))
require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "ca-name", ""))
assert.Contains(t, fw.InRules.UDP[1].CANames, "ca-name")
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
require.NoError(t, fw.AddRule(true, firewall.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "", "ca-sha"))
require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "", "ca-sha"))
assert.Contains(t, fw.InRules.UDP[1].CAShas, "ca-sha")
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
@@ -185,7 +186,7 @@ func TestFirewall_Drop(t *testing.T) {
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
LocalPort: 10,
RemotePort: 90,
Protocol: firewall.ProtoUDP,
Protocol: iputil.IPProtocolUDP,
Fragment: false,
}
@@ -263,7 +264,7 @@ func TestFirewall_DropV6(t *testing.T) {
RemoteAddr: netip.MustParseAddr("fd12::34"),
LocalPort: 10,
RemotePort: 90,
Protocol: firewall.ProtoUDP,
Protocol: iputil.IPProtocolUDP,
Fragment: false,
}
@@ -350,7 +351,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
Certificate: &dummyCert{},
}
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoUDP}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolUDP}, true, c, cp))
}
})
@@ -360,7 +361,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
Certificate: &dummyCert{},
}
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 1}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 1}, true, c, cp))
}
})
@@ -370,7 +371,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
}
ip := netip.MustParsePrefix("9.254.254.254/32")
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
}
})
b.Run("pass proto, port, fail on local CIDRv6", func(b *testing.B) {
@@ -379,7 +380,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
}
ip := netip.MustParsePrefix("fd99::99/128")
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
}
})
@@ -392,7 +393,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"nope": {}},
}
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
}
})
b.Run("pass proto, port, any local CIDRv6, fail all group, name, and cidr", func(b *testing.B) {
@@ -404,7 +405,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"nope": {}},
}
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
}
})
@@ -417,7 +418,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"nope": {}},
}
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
}
})
b.Run("pass proto, port, specific local CIDRv6, fail all group, name, and cidr", func(b *testing.B) {
@@ -429,7 +430,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"nope": {}},
}
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
}
})
@@ -441,7 +442,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"good-group": {}},
}
for n := 0; n < b.N; n++ {
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp))
assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
}
})
@@ -453,7 +454,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"good-group": {}},
}
for n := 0; n < b.N; n++ {
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
}
})
b.Run("pass on group on specific local cidr6", func(b *testing.B) {
@@ -464,7 +465,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"good-group": {}},
}
for n := 0; n < b.N; n++ {
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
}
})
@@ -476,7 +477,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"nope": {}},
}
for n := 0; n < b.N; n++ {
ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp)
ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp)
}
})
}
@@ -492,7 +493,7 @@ func TestFirewall_Drop2(t *testing.T) {
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
LocalPort: 10,
RemotePort: 90,
Protocol: firewall.ProtoUDP,
Protocol: iputil.IPProtocolUDP,
Fragment: false,
}
@@ -550,7 +551,7 @@ func TestFirewall_Drop3(t *testing.T) {
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
LocalPort: 1,
RemotePort: 1,
Protocol: firewall.ProtoUDP,
Protocol: iputil.IPProtocolUDP,
Fragment: false,
}
@@ -638,7 +639,7 @@ func TestFirewall_Drop3V6(t *testing.T) {
RemoteAddr: netip.MustParseAddr("fd12::34"),
LocalPort: 1,
RemotePort: 1,
Protocol: firewall.ProtoUDP,
Protocol: iputil.IPProtocolUDP,
Fragment: false,
}
@@ -675,7 +676,7 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
LocalPort: 10,
RemotePort: 90,
Protocol: firewall.ProtoUDP,
Protocol: iputil.IPProtocolUDP,
Fragment: false,
}
network := netip.MustParsePrefix("1.2.3.4/24")
@@ -758,13 +759,13 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
templ := firewall.Packet{
LocalAddr: netip.MustParseAddr("1.2.3.4"),
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
Protocol: firewall.ProtoICMP,
Protocol: iputil.IPProtocolICMP,
Fragment: false,
}
t.Run("ICMP allowed", func(t *testing.T) {
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
require.NoError(t, fw.AddRule(true, firewall.ProtoICMP, 0, 0, []string{"any"}, "", "", "", "", ""))
require.NoError(t, fw.AddRule(true, iputil.IPProtocolICMP, 0, 0, []string{"any"}, "", "", "", "", ""))
t.Run("zero ports", func(t *testing.T) {
p := templ.Copy()
p.LocalPort = 0
@@ -910,7 +911,7 @@ func TestFirewall_DropIPSpoofing(t *testing.T) {
RemoteAddr: netip.MustParseAddr("192.0.2.3"),
LocalPort: 1,
RemotePort: 1,
Protocol: firewall.ProtoUDP,
Protocol: iputil.IPProtocolUDP,
Fragment: false,
}
assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP)
@@ -961,7 +962,7 @@ func TestFirewall_ConntrackSourceSpoofingAcrossPeers(t *testing.T) {
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
LocalPort: 443,
RemotePort: 55000,
Protocol: firewall.ProtoUDP,
Protocol: iputil.IPProtocolUDP,
}
require.NoError(t, fw.Drop(flow, true, &victimHI, cp, nil),
@@ -1031,7 +1032,7 @@ func BenchmarkFirewallDropConntrackHit(b *testing.B) {
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
LocalPort: 443,
RemotePort: 55000,
Protocol: firewall.ProtoUDP,
Protocol: iputil.IPProtocolUDP,
}
cases := []struct {
@@ -1317,28 +1318,28 @@ func TestAddFirewallRulesFromConfig(t *testing.T) {
mf := &mockFirewall{}
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "tcp", "host": "a"}}}
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoTCP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolTCP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
// Test adding udp rule
conf = config.NewC(test.NewLogger())
mf = &mockFirewall{}
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "udp", "host": "a"}}}
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoUDP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolUDP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
// Test adding icmp rule
conf = config.NewC(test.NewLogger())
mf = &mockFirewall{}
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "icmp", "host": "a"}}}
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
// Test adding icmp rule no port
conf = config.NewC(test.NewLogger())
mf = &mockFirewall{}
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"proto": "icmp", "host": "a"}}}
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
// Test adding any rule
conf = config.NewC(test.NewLogger())
@@ -1582,7 +1583,7 @@ func buildTestCase(setup testsetup, err error, theirPrefixes ...netip.Prefix) te
RemoteAddr: theirPrefixes[0].Addr(),
LocalPort: 10,
RemotePort: 90,
Protocol: firewall.ProtoUDP,
Protocol: iputil.IPProtocolUDP,
Fragment: false,
}
return testcase{
+3 -1
View File
@@ -12,7 +12,7 @@ require (
github.com/gogo/protobuf v1.3.2
github.com/google/gopacket v1.1.19
github.com/kardianos/service v1.3.0
github.com/miekg/dns v1.1.73
github.com/miekg/dns v1.1.72
github.com/miekg/pkcs11 v1.1.2
github.com/nbrownus/go-metrics-prometheus v0.0.0-20210712211119-974a6260965f
github.com/prometheus/client_golang v1.24.1
@@ -45,5 +45,7 @@ require (
github.com/prometheus/common v0.70.1 // indirect
github.com/prometheus/procfs v0.21.1 // indirect
github.com/vishvananda/netns v0.0.5 // indirect
golang.org/x/mod v0.36.0 // indirect
golang.org/x/time v0.5.0 // indirect
golang.org/x/tools v0.45.0 // indirect
)
+6 -2
View File
@@ -81,8 +81,8 @@ github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI=
github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc=
github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw=
github.com/matttproud/golang_protobuf_extensions v1.0.1/go.mod h1:D8He9yQNgCq6Z5Ld7szi9bcBfOoFv/3dc6xSMkL2PC0=
github.com/miekg/dns v1.1.73 h1:uhT8nJxmTrPJYClxVxTCX+CVn6qnzSiybRk72Z6DgrE=
github.com/miekg/dns v1.1.73/go.mod h1:RW2Obtfd5NZHvOFe3zYG0W8koWOQtAzyHaLo8vASBuQ=
github.com/miekg/dns v1.1.72 h1:vhmr+TF2A3tuoGNkLDFK9zi36F2LS+hKTRW0Uf8kbzI=
github.com/miekg/dns v1.1.72/go.mod h1:+EuEPhdHOsfk6Wk5TT2CzssZdqkmFhf8r+aVyDEToIs=
github.com/miekg/pkcs11 v1.1.2 h1:/VxmeAX5qU6Q3EwafypogwWbYryHFmF2RpkJmw3m4MQ=
github.com/miekg/pkcs11 v1.1.2/go.mod h1:XsNlhZGX73bx86s2hdc/FuaLm2CPZJemRLMA+WTFxgs=
github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q=
@@ -161,6 +161,8 @@ golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPI
golang.org/x/mod v0.1.1-0.20191105210325-c90efee705ee/go.mod h1:QqPTAvyqsEbceGzBzNggFXnrqF1CaUcvgkdR5Ot7KZg=
golang.org/x/mod v0.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
golang.org/x/mod v0.36.0 h1:JJjpVx6myfUsUdAzZuOSTTmRE0PfZeNWzzvKrP7amb4=
golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ=
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20181114220301-adae6a3d119a/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20190108225652-1e06a53dbb7e/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
@@ -212,6 +214,8 @@ golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtn
golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
golang.org/x/tools v0.0.0-20200619180055-7c47624df98f/go.mod h1:EkVYQZoAsY45+roYkvgYkIh4xh/qjgUK9TdY2XT94GE=
golang.org/x/tools v0.0.0-20210106214847-113979e3529a/go.mod h1:emZCQorbCU4vsT4fOWvOPXz4eW1wZW4PmDk9uLelYpA=
golang.org/x/tools v0.45.0 h1:18qN3FAooORvApf5XjCXgsuayZOEtXf6JK18I3+ONa8=
golang.org/x/tools v0.45.0/go.mod h1:LuUGqqaXcXMEFEruIVJVm5mgDD8vww/z/SR1gQ4uE/0=
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
+67 -13
View File
@@ -239,11 +239,15 @@ const (
type HostInfo struct {
remote atomic.Pointer[netip.AddrPort]
remotes *RemoteList
promoteCounter atomic.Uint32
ConnectionState *ConnectionState
remoteIndexId uint32
localIndexId uint32
// 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
@@ -262,11 +266,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 +274,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.
@@ -669,7 +665,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)
}
@@ -770,6 +766,64 @@ 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
// The epoch is the top 29 bits, it would take 2^29 rebinds to wrap and we will never get there
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)
}
}
// 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
+46
View File
@@ -401,3 +401,49 @@ func TestHostMap_RelayState(t *testing.T) {
assert.Equal(t, []netip.Addr{}, h1.relayState.relays)
}
// sentSinceCheck reports whether anything has been sent since the connection manager last looked. Test only:
// production reads the out bit through takeTraffic on the connection manager tick.
func (i *HostInfo) sentSinceCheck() bool {
return i.state.Load()&stateOut != 0
}
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))
}
+14 -18
View File
@@ -58,6 +58,9 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Parse
// kernel as one giant blob; segment first so the loopback
// path sees one IP datagram per Write.
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
// The kernel may have left the transport checksum for hardware
// offload to finish; nothing between here and the tun will.
iputil.SetTransportChecksum(seg)
_, werr := f.queues[q].Write(seg)
return werr
})
@@ -160,21 +163,18 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
// One traffic-out mark covers every segment of the superpacket; doing it
// per segment in sendInsideEncrypt paid an atomic store up to ~45 extra
// times per TSO packet, inside writeLock under boring crypto.
f.connectionManager.Out(hostinfo)
remote := hostinfo.GetRemote()
if 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) {
f.lightHouse.QueryServer(hostinfo.vpnAddrs[0])
hostinfo.lastRebindCount = f.rebindCount
if f.l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(f.l).Debug("Lighthouse update triggered for punch due to rebind counter",
hostinfo.logger(f.l).Debug("Lighthouse update triggered for punch due to rebind epoch",
"vpnAddrs", hostinfo.vpnAddrs,
)
}
}
remote := hostinfo.GetRemote()
if !remote.IsValid() { //the relay path
//first, find our relay hostinfo:
var relayHostInfo *HostInfo
@@ -460,7 +460,7 @@ func (f *Interface) prepareSendVia(via *HostInfo,
}
out = header.Encode(out, header.Version, header.Message, header.MessageRelay, relay.RemoteIndex, c)
f.connectionManager.Out(via)
f.connectionManager.OutNoRebind(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.
@@ -553,17 +553,13 @@ 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.
// A closing tunnel is torn down right after this, so skip the connection manager entirely: no point recording
// traffic or asking the lighthouse for a punch. Otherwise, if we rebound since this tunnel last sent, ask the
// lighthouse to get the far side punching at us again.
if t != header.CloseTunnel && f.connectionManager.Out(hostinfo) {
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,
)
}
+265
View File
@@ -0,0 +1,265 @@
package nebula
import (
"encoding/binary"
"io"
"net/netip"
"testing"
"github.com/gaissmai/bart"
"github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/iputil"
"github.com/slackhq/nebula/overlay/tio"
"github.com/slackhq/nebula/test"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
const (
ipv4HeaderLen = 20
ipv6HeaderLen = 40
)
// capturingTun is a tio.Queue that records what is written to it. A queue that
// discards writes is indistinguishable from a packet that was never forwarded.
type capturingTun struct {
writes [][]byte
}
func (c *capturingTun) Read() ([]tio.Packet, error) { return nil, io.EOF }
func (c *capturingTun) Close() error { return nil }
func (c *capturingTun) Write(b []byte) (int, error) {
c.writes = append(c.writes, append([]byte(nil), b...))
return len(b), nil
}
func newSelfForwardInterface(myAddrs ...netip.Addr) (*Interface, *capturingTun) {
vpnAddrs := &bart.Lite{}
for _, a := range myAddrs {
vpnAddrs.Insert(netip.PrefixFrom(a, a.BitLen()))
}
tun := &capturingTun{}
return &Interface{
l: test.NewLogger(),
myVpnAddrsTable: vpnAddrs,
myBroadcastAddrsTable: &bart.Lite{},
queues: []tio.Queue{tun},
}, tun
}
func consumeInside(f *Interface, packet []byte) {
f.consumeInsidePacket(tio.Packet{Bytes: packet}, &firewall.ParsedPacket{}, make([]byte, 12), nil, make([]byte, mtu), 0, nil)
}
// l4Proto describes one upper-layer header for these tests: its IP next-header
// value, where its checksum field sits within the header, and how to build a
// minimal instance of it.
type l4Proto struct {
name string
nextHdr uint8
cksumAt int
build func() []byte
}
var (
tcpSyn = l4Proto{"tcp", iputil.IPProtocolTCP, 16, func() []byte {
h := make([]byte, 20)
binary.BigEndian.PutUint16(h[0:2], 49152)
binary.BigEndian.PutUint16(h[2:4], 443)
binary.BigEndian.PutUint32(h[4:8], 0x11223344) // sequence
h[12] = 5 << 4 // data offset, no options
h[13] = 0x02 // SYN
binary.BigEndian.PutUint16(h[14:16], 65535) // window
return h
}}
udpDatagram = l4Proto{"udp", iputil.IPProtocolUDP, 6, func() []byte {
h := make([]byte, 8+4)
binary.BigEndian.PutUint16(h[0:2], 49152)
binary.BigEndian.PutUint16(h[2:4], 53)
binary.BigEndian.PutUint16(h[4:6], uint16(len(h)))
copy(h[8:], "ping")
return h
}}
icmpEcho = l4Proto{"icmp", iputil.IPProtocolICMP, 2, func() []byte { return echoRequest(8) }}
icmpv6Echo = l4Proto{"icmpv6", iputil.IPProtocolICMPv6, 2, func() []byte { return echoRequest(128) }}
)
// echoRequest builds an echo request body. The type differs between ICMP and
// ICMPv6, the rest of the header does not.
func echoRequest(typ uint8) []byte {
h := make([]byte, 8)
h[0] = typ
binary.BigEndian.PutUint16(h[4:6], 0xbeef) // identifier
binary.BigEndian.PutUint16(h[6:8], 1) // sequence
return h
}
func buildIPv6(src, dst netip.Addr, p l4Proto) []byte {
l4 := p.build()
pkt := make([]byte, ipv6HeaderLen+len(l4))
pkt[0] = 0x60
binary.BigEndian.PutUint16(pkt[4:6], uint16(len(l4)))
pkt[6] = p.nextHdr
pkt[7] = 64
copy(pkt[8:24], src.AsSlice())
copy(pkt[24:40], dst.AsSlice())
copy(pkt[ipv6HeaderLen:], l4)
if l4 := pkt[ipv6HeaderLen:]; p.nextHdr == iputil.IPProtocolTCP || p.nextHdr == iputil.IPProtocolUDP {
sum := ipv6PseudoheaderSum(src, dst, uint32(p.nextHdr), uint32(len(l4)))
binary.BigEndian.PutUint16(l4[p.cksumAt:], ^fold(sumBytes(l4, sum)))
}
return pkt
}
func buildIPv4(src, dst netip.Addr, p l4Proto) []byte {
l4 := p.build()
pkt := make([]byte, ipv4HeaderLen+len(l4))
pkt[0] = 0x45
binary.BigEndian.PutUint16(pkt[2:4], uint16(len(pkt)))
pkt[8] = 64
pkt[9] = p.nextHdr
copy(pkt[12:16], src.AsSlice())
copy(pkt[16:20], dst.AsSlice())
copy(pkt[ipv4HeaderLen:], l4)
if l4 := pkt[ipv4HeaderLen:]; p.nextHdr == iputil.IPProtocolTCP || p.nextHdr == iputil.IPProtocolUDP {
sum := sumBytes(pkt[12:20], uint32(p.nextHdr)+uint32(len(l4)))
binary.BigEndian.PutUint16(l4[p.cksumAt:], ^fold(sumBytes(l4, sum)))
}
return pkt
}
// ipv6PseudoheaderSum is the RFC 2460 section 8.1 pseudo-header sum: source,
// destination, a 32 bit upper-layer packet length and a 32 bit zero-padded next
// header. Kept local to the test so these assertions do not check nebula's
// checksum code against itself.
func ipv6PseudoheaderSum(src, dst netip.Addr, nextHeader, length uint32) uint32 {
var csum uint32
s, d := src.AsSlice(), dst.AsSlice()
for i := 0; i < 16; i += 2 {
csum += uint32(s[i])<<8 | uint32(s[i+1])
csum += uint32(d[i])<<8 | uint32(d[i+1])
}
return csum + length + nextHeader
}
func sumBytes(b []byte, csum uint32) uint32 {
for i := 0; i+1 < len(b); i += 2 {
csum += uint32(b[i])<<8 | uint32(b[i+1])
}
if len(b)%2 == 1 {
csum += uint32(b[len(b)-1]) << 8
}
return csum
}
func fold(csum uint32) uint16 {
for csum > 0xffff {
csum = (csum >> 16) + (csum & 0xffff)
}
return uint16(csum)
}
// l4ChecksumValid6 verifies an IPv6 upper-layer checksum the way a receiver
// does: the pseudo-header plus the whole upper-layer segment, checksum field
// included, folds to 0xffff. The next header field is the upper-layer protocol
// only while there are no extension headers, which is all this file builds.
func l4ChecksumValid6(pkt []byte) bool {
src, _ := netip.AddrFromSlice(pkt[8:24])
dst, _ := netip.AddrFromSlice(pkt[24:40])
l4 := pkt[ipv6HeaderLen:]
return fold(sumBytes(l4, ipv6PseudoheaderSum(src, dst, uint32(pkt[6]), uint32(len(l4))))) == 0xffff
}
// l4ChecksumValid4 is the IPv4 counterpart: the RFC 793/768 pseudo-header is
// source, destination, a zero byte, the protocol and the upper-layer length.
func l4ChecksumValid4(pkt []byte) bool {
ihl := int(pkt[0]&0x0f) << 2
l4 := pkt[ihl:]
return fold(sumBytes(l4, sumBytes(pkt[12:20], uint32(pkt[9])+uint32(len(l4))))) == 0xffff
}
// TestConsumeInsidePacketSelfTraffic covers the self-addressed branch of
// consumeInsidePacket, taken where immediatelyForwardToSelf is set (see
// inside_bsd.go): the packet goes straight back to the tun, ahead of the
// firewall and the handshake.
func TestConsumeInsidePacketSelfTraffic(t *testing.T) {
v4 := netip.MustParseAddr("100.100.1.42")
v6 := netip.MustParseAddr("fd00::42")
tests := []struct {
name string
addr netip.Addr
pkt []byte
}{
{"ipv4/tcp", v4, buildIPv4(v4, v4, tcpSyn)},
{"ipv4/udp", v4, buildIPv4(v4, v4, udpDatagram)},
{"ipv4/icmp", v4, buildIPv4(v4, v4, icmpEcho)},
{"ipv6/tcp", v6, buildIPv6(v6, v6, tcpSyn)},
{"ipv6/udp", v6, buildIPv6(v6, v6, udpDatagram)},
{"ipv6/icmpv6", v6, buildIPv6(v6, v6, icmpv6Echo)},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
f, tun := newSelfForwardInterface(tt.addr)
// consumeInsidePacket writes through the slice it is handed, so a
// packet that arrived with a valid checksum must come back out of
// bytes taken before the call, unchanged.
want := append([]byte(nil), tt.pkt...)
consumeInside(f, tt.pkt)
if immediatelyForwardToSelf {
require.Len(t, tun.writes, 1)
assert.Equal(t, want, tun.writes[0])
} else {
assert.Empty(t, tun.writes, "self traffic reaches the tun over loopback here and must be dropped")
}
})
}
}
// TestConsumeInsidePacketSelfTrafficChecksum shows that the self-forward
// returns the bytes it was handed, so a packet that arrived with a wrong
// upper-layer checksum is written back with that same wrong checksum and the
// kernel drops it on re-entry.
//
// This is how a macOS host loses TCP and UDP to its own IPv6 overlay address:
// the kernel writes only the pseudo-header sum into the checksum field and
// defers completion to hardware offload, state that does not survive the
// crossing into userspace. Which kernels do this, for which protocols and IP
// versions, is a property of the kernel and belongs to a test against a live
// one; here the checksum is simply wrong, and the forward must make it right.
func TestConsumeInsidePacketSelfTrafficChecksum(t *testing.T) {
if !immediatelyForwardToSelf {
t.Skip("self traffic never reaches the tun on this platform")
}
versions := []struct {
name string
addr netip.Addr
build func(src, dst netip.Addr, p l4Proto) []byte
l4At int
valid func(pkt []byte) bool
}{
{"v4", netip.MustParseAddr("100.100.1.42"), buildIPv4, ipv4HeaderLen, l4ChecksumValid4},
{"v6", netip.MustParseAddr("fd00::42"), buildIPv6, ipv6HeaderLen, l4ChecksumValid6},
}
for _, v := range versions {
for _, p := range []l4Proto{tcpSyn, udpDatagram} {
t.Run(v.name+"/"+p.name, func(t *testing.T) {
pkt := v.build(v.addr, v.addr, p)
binary.BigEndian.PutUint16(pkt[v.l4At+p.cksumAt:], 0x1234)
require.False(t, v.valid(pkt), "the packet under test must start with a wrong checksum")
f, tun := newSelfForwardInterface(v.addr)
consumeInside(f, pkt)
require.Len(t, tun.writes, 1)
assert.True(t, v.valid(tun.writes[0]),
"a forwarded %s packet must carry a valid checksum, got 0x%04x",
p.name, binary.BigEndian.Uint16(tun.writes[0][v.l4At+p.cksumAt:]))
})
}
}
}
+2 -2
View File
@@ -107,8 +107,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
+146
View File
@@ -0,0 +1,146 @@
package iputil
import (
"encoding/binary"
"github.com/slackhq/nebula/overlay/checksum"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
)
const udpHeaderLen = 8
// SetTransportChecksum recomputes the TCP or UDP checksum of an IPv4 or IPv6
// packet in place.
//
// A kernel that offloads checksums to the NIC hands a packet to a tun with the
// transport checksum unfinished: only the pseudo-header sum is in the field and
// the rest is left for hardware that a tun does not have. A packet written
// straight back to that tun is dropped on re-entry unless the checksum is
// completed first. ICMP is left alone; it arrived complete on the kernels this
// was measured against.
//
// So is any packet whose transport header cannot be located: fragments, unknown
// extension headers and truncated packets. An IPv6 fragment header is declined
// even when it carries the whole datagram (RFC 6946 atomic fragment), because
// the walk reports only that a fragment header was present.
func SetTransportChecksum(packet []byte) {
if len(packet) < 1 {
return
}
switch int(packet[0] >> 4) {
case ipv4.Version:
setTransportChecksum4(packet)
case ipv6.Version:
setTransportChecksum6(packet)
}
}
func setTransportChecksum4(packet []byte) {
if len(packet) < ipv4.HeaderLen {
return
}
ihl := int(packet[0]&0x0f) << 2
end := int(binary.BigEndian.Uint16(packet[2:4]))
if ihl < ipv4.HeaderLen || end < ihl || end > len(packet) {
return
}
// The checksum covers the whole datagram, which a fragment (MF set or a
// non-zero offset) does not carry.
if binary.BigEndian.Uint16(packet[6:8])&0x3fff != 0 {
return
}
transport, ok := transportExtent(packet[ihl:end], packet[9])
if !ok {
return
}
csum := ipv4PseudoheaderChecksum(packet[12:16], packet[16:20], uint32(packet[9]), uint32(len(transport)))
writeTransportChecksum(transport, packet[9], csum)
}
func setTransportChecksum6(packet []byte) {
if len(packet) < ipv6.HeaderLen {
return
}
end := ipv6.HeaderLen + int(binary.BigEndian.Uint16(packet[4:6]))
if end > len(packet) {
return
}
// The checksum covers the whole datagram, which a fragment does not carry.
// An unknown extension header hides where the transport header starts. A
// chain longer than the walk's budget ends it early, at an offset that was
// never checked against the packet.
proto, offset, _, anyFragment, err := IPv6FindUpperProtocol(packet[:end])
if err != nil || anyFragment || offset >= end {
return
}
transport, ok := transportExtent(packet[offset:end], proto)
if !ok {
return
}
csum := ipv6PseudoheaderChecksum(packet[8:24], packet[24:40], uint32(proto), uint32(len(transport)))
writeTransportChecksum(transport, proto, csum)
}
// transportExtent narrows a segment to the length its own header declares. UDP
// carries a Length field, and RFC 768 and RFC 8200 section 8.1 both make that
// field, not the IP payload extent, the length the pseudo-header counts and the
// checksum covers; a datagram padded out to a link's minimum frame is the usual
// way the two differ. TCP has no such field, so its segment runs to the end of
// the IP payload. A Length that overruns the bytes IP delivered describes a
// datagram that is not there.
func transportExtent(transport []byte, proto uint8) ([]byte, bool) {
if proto != IPProtocolUDP {
return transport, true
}
if len(transport) < udpHeaderLen {
return nil, false
}
ulen := int(binary.BigEndian.Uint16(transport[4:6]))
if ulen < udpHeaderLen || ulen > len(transport) {
return nil, false
}
return transport[:ulen], true
}
// writeTransportChecksum stores the checksum of transport, taken over the
// pseudo-header sum csum, in the header's checksum field. A UDP checksum that
// computes to zero goes on the wire as 0xffff: zero means no checksum was
// computed (RFC 768), and over IPv6 the checksum is mandatory (RFC 8200
// section 8.1).
func writeTransportChecksum(transport []byte, proto uint8, csum uint32) {
var at, minLen int
switch proto {
case IPProtocolTCP:
at, minLen = 16, 20
case IPProtocolUDP:
at, minLen = 6, udpHeaderLen
default:
return
}
if len(transport) < minLen {
return
}
transport[at], transport[at+1] = 0, 0
sum := ^checksum.Checksum(transport, fold(csum))
if sum == 0 && proto == IPProtocolUDP {
sum = 0xffff
}
binary.BigEndian.PutUint16(transport[at:], sum)
}
// fold reduces a pseudo-header sum to the 16 bit seed Checksum takes. Carrying
// the high half back into the low half is what keeps the reduction lossless, so
// the seed sums exactly as the wider value would; 0xffff is its fixed point.
// Every term of that sum comes from a 16 bit field, so it stays far below the
// width at which the accumulator would wrap.
func fold(csum uint32) uint16 {
for csum > 0xffff {
csum = (csum >> 16) + (csum & 0xffff)
}
return uint16(csum)
}
+242
View File
@@ -0,0 +1,242 @@
package iputil
import (
"encoding/binary"
"net"
"testing"
"github.com/google/gopacket"
"github.com/google/gopacket/layers"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/net/ipv6"
)
// serialize builds a packet with gopacket, whose checksums are computed
// independently of this package.
func serialize(t *testing.T, ls ...gopacket.SerializableLayer) []byte {
buf := gopacket.NewSerializeBuffer()
require.NoError(t, gopacket.SerializeLayers(buf, gopacket.SerializeOptions{FixLengths: true, ComputeChecksums: true}, ls...))
return append([]byte(nil), buf.Bytes()...)
}
// withExtensionHeader inserts an 8 byte IPv6 extension header of the given
// type between the IPv6 header and its payload. The transport checksum does not
// change: the pseudo-header counts only upper-layer bytes.
func withExtensionHeader(pkt []byte, typ layers.IPProtocol, hdr [8]byte) []byte {
hdr[0] = pkt[6]
out := make([]byte, 0, len(pkt)+8)
out = append(out, pkt[:40]...)
out = append(out, hdr[:]...)
out = append(out, pkt[40:]...)
out[6] = byte(typ)
binary.BigEndian.PutUint16(out[4:6], binary.BigEndian.Uint16(pkt[4:6])+8)
return out
}
// truncate copies the first n bytes into a buffer of exactly that capacity, so
// a read past the length panics instead of quietly succeeding.
func truncate(pkt []byte, n int) []byte {
out := make([]byte, n)
copy(out, pkt)
return out
}
// extChain builds an IPv6 packet fronted by n Destination Options headers. Each
// points at another one, so the walk spends its whole budget without reaching a
// transport header. lastExtLen inflates the final header's declared length,
// which is how the walk ends up past the end of the packet.
func extChain(n int, lastExtLen byte) []byte {
pkt := make([]byte, ipv6.HeaderLen)
pkt[0], pkt[6], pkt[7] = 0x60, 60, 64
for i := range n {
h := make([]byte, 8)
h[0] = 60
if i == n-1 {
h[1] = lastExtLen
}
pkt = append(pkt, h...)
}
pkt = append(pkt, make([]byte, 20)...)
binary.BigEndian.PutUint16(pkt[4:6], uint16(len(pkt)-ipv6.HeaderLen))
return pkt
}
func TestSetTransportChecksum(t *testing.T) {
// Source and destination differ so that a pseudo-header built from the wrong
// one, or from the two swapped, does not land on the same checksum anyway.
v4 := func(proto layers.IPProtocol) *layers.IPv4 {
return &layers.IPv4{Version: 4, TTL: 64, Id: 0x1234, Protocol: proto, SrcIP: net.IPv4(192, 0, 2, 1).To4(), DstIP: net.IPv4(198, 51, 100, 2).To4()}
}
v6 := func(proto layers.IPProtocol) *layers.IPv6 {
return &layers.IPv6{Version: 6, HopLimit: 64, NextHeader: proto, SrcIP: net.ParseIP("2001:db8::1"), DstIP: net.ParseIP("2001:db8:1::2")}
}
tcp := func(ip gopacket.NetworkLayer) *layers.TCP {
l := &layers.TCP{SrcPort: 49152, DstPort: 443, SYN: true, Window: 65535}
require.NoError(t, l.SetNetworkLayerForChecksum(ip))
return l
}
udp := func(ip gopacket.NetworkLayer) *layers.UDP {
l := &layers.UDP{SrcPort: 49152, DstPort: 53}
require.NoError(t, l.SetNetworkLayerForChecksum(ip))
return l
}
payload := gopacket.Payload("self")
nop := layers.IPv4Option{OptionType: 1, OptionLength: 1}
ip4tcp := v4(layers.IPProtocolTCP)
ip4opts := v4(layers.IPProtocolTCP)
ip4opts.Options = []layers.IPv4Option{nop, nop, nop, nop}
ip4udp := v4(layers.IPProtocolUDP)
ip6tcp := v6(layers.IPProtocolTCP)
ip6udp := v6(layers.IPProtocolUDP)
hopByHop := [8]byte{0, 0, 1, 4} // next header, length 0, PadN of 4
// Bytes past the length the IP header declares are not part of the
// datagram and must not be summed.
trailing4 := append(serialize(t, ip4tcp, tcp(ip4tcp), payload), []byte("trailing")...)
trailing6 := append(serialize(t, ip6tcp, tcp(ip6tcp), payload), []byte("trailing")...)
// A datagram padded out past the length UDP declares: the pseudo-header
// counts the UDP Length field, so the checksum is the unpadded one.
padded4 := append(serialize(t, ip4udp, udp(ip4udp), payload), []byte("pad!")...)
binary.BigEndian.PutUint16(padded4[2:4], uint16(len(padded4)))
padded6 := append(serialize(t, ip6udp, udp(ip6udp), payload), []byte("pad!")...)
binary.BigEndian.PutUint16(padded6[4:6], uint16(len(padded6)-ipv6.HeaderLen))
// Corrupting the checksum and asking for it back must yield gopacket's
// packet, byte for byte.
recomputed := []struct {
name string
pkt []byte
cksum int
}{
{"v4 tcp", serialize(t, ip4tcp, tcp(ip4tcp), payload), 20 + 16},
{"v4 tcp with ip options", serialize(t, ip4opts, tcp(ip4opts), payload), 24 + 16},
{"v4 udp", serialize(t, ip4udp, udp(ip4udp), payload), 20 + 6},
{"v4 tcp header only", serialize(t, ip4tcp, tcp(ip4tcp)), 20 + 16},
{"v4 udp header only", serialize(t, ip4udp, udp(ip4udp)), 20 + 6},
{"v6 tcp", serialize(t, ip6tcp, tcp(ip6tcp), payload), 40 + 16},
{"v6 udp", serialize(t, ip6udp, udp(ip6udp), payload), 40 + 6},
{"v6 udp header only", serialize(t, ip6udp, udp(ip6udp)), 40 + 6},
{"v6 tcp behind hop-by-hop", withExtensionHeader(serialize(t, ip6tcp, tcp(ip6tcp), payload), layers.IPProtocolIPv6HopByHop, hopByHop), 48 + 16},
{"v4 tcp with bytes past the total length", trailing4, 20 + 16},
{"v6 tcp with bytes past the payload length", trailing6, 40 + 16},
{"v4 udp padded past its declared length", padded4, 20 + 6},
{"v6 udp padded past its declared length", padded6, 40 + 6},
}
for _, tt := range recomputed {
t.Run(tt.name, func(t *testing.T) {
got := append([]byte(nil), tt.pkt...)
binary.BigEndian.PutUint16(got[tt.cksum:], 0x1234)
require.NotEqual(t, tt.pkt, got)
SetTransportChecksum(got)
assert.Equal(t, tt.pkt, got)
})
}
ip4frag := v4(layers.IPProtocolTCP)
ip4frag.Flags = layers.IPv4MoreFragments
ip4later := v4(layers.IPProtocolTCP)
ip4later.FragOffset = 1
ip4icmp := v4(layers.IPProtocolICMPv4)
badIHL := serialize(t, ip4tcp, tcp(ip4tcp), payload)
badIHL[0] = 0x44 // header length 16, shorter than an ipv4 header
shortTotalLen := serialize(t, ip4tcp, tcp(ip4tcp), payload)
binary.BigEndian.PutUint16(shortTotalLen[2:4], 10) // shorter than the header it introduces
cutTCP := serialize(t, ip4tcp, tcp(ip4tcp), payload)
binary.BigEndian.PutUint16(cutTCP[2:4], 20+19) // one byte short of a tcp header
cutTCP = truncate(cutTCP, 20+19)
cutUDP := serialize(t, ip4udp, udp(ip4udp), payload)
binary.BigEndian.PutUint16(cutUDP[2:4], 20+7) // one byte short of a udp header
cutUDP = truncate(cutUDP, 20+7)
// Two bytes short, so a transport header survives whole and the minimum
// length check cannot stand in for the bounds check.
cutV6 := truncate(serialize(t, ip6tcp, tcp(ip6tcp), payload), 62)
fragment := [8]byte{0, 0, 0, 1, 0, 0, 0, 1} // next header, reserved, offset 0 with M set, id
overrun4 := serialize(t, ip4udp, udp(ip4udp), payload)
binary.BigEndian.PutUint16(overrun4[24:26], uint16(len(overrun4)-20+1)) // one byte past what ip delivered
overrun6 := serialize(t, ip6udp, udp(ip6udp), payload)
binary.BigEndian.PutUint16(overrun6[44:46], uint16(len(overrun6)-ipv6.HeaderLen+1))
shortUDPLen := serialize(t, ip4udp, udp(ip4udp), payload)
binary.BigEndian.PutUint16(shortUDPLen[24:26], 7) // shorter than the header it counts
// Where the checksum cannot be completed the packet is left as it came.
untouched := []struct {
name string
pkt []byte
cksum int
}{
{"v4 first fragment", serialize(t, ip4frag, tcp(ip4frag), payload), 20 + 16},
{"v4 later fragment", serialize(t, ip4later, tcp(ip4later), payload), 20 + 16},
{"v4 icmp", serialize(t, ip4icmp, &layers.ICMPv4{TypeCode: layers.CreateICMPv4TypeCode(8, 0), Id: 1, Seq: 1}, payload), 20 + 2},
{"v4 header length below the minimum", badIHL, 20 + 16},
{"v4 total length below the header length", shortTotalLen, 20 + 16},
{"v4 truncated below its total length", truncate(serialize(t, ip4tcp, tcp(ip4tcp), payload), 30), -1},
{"v4 tcp header cut short", cutTCP, 20 + 16},
{"v4 udp header cut short", cutUDP, -1},
{"v6 fragment", withExtensionHeader(serialize(t, ip6tcp, tcp(ip6tcp), payload), layers.IPProtocolIPv6Fragment, fragment), 48 + 16},
{"v6 truncated below its payload length", truncate(serialize(t, ip6tcp, tcp(ip6tcp), payload), 50), -1},
{"v6 truncated with a whole transport header still present", cutV6, 40 + 16},
{"v6 extension header chain longer than the walk", extChain(9, 0), 112 + 16},
{"v6 extension header chain running past the packet", extChain(8, 255), 104 + 16},
{"v4 udp length past the end of the datagram", overrun4, 20 + 6},
{"v6 udp length past the end of the datagram", overrun6, 40 + 6},
{"v4 udp length below a udp header", shortUDPLen, 20 + 6},
}
for _, tt := range untouched {
t.Run(tt.name, func(t *testing.T) {
if tt.cksum >= 0 {
binary.BigEndian.PutUint16(tt.pkt[tt.cksum:], 0x1234)
}
want := append([]byte(nil), tt.pkt...)
SetTransportChecksum(tt.pkt)
assert.Equal(t, want, tt.pkt)
})
}
t.Run("too short to carry a header", func(t *testing.T) {
for _, pkt := range [][]byte{nil, {}, {0x45}, {0x60}} {
assert.NotPanics(t, func() { SetTransportChecksum(pkt) })
}
})
t.Run("tcp checksum of zero goes out as zero", func(t *testing.T) {
pkt := serialize(t, ip4tcp, tcp(ip4tcp), gopacket.Payload{0, 0})
c := binary.BigEndian.Uint16(pkt[36:38])
require.NotZero(t, c)
// Only udp reserves zero to mean "not computed", so tcp keeps it.
binary.BigEndian.PutUint16(pkt[40:42], c)
SetTransportChecksum(pkt)
assert.Zero(t, binary.BigEndian.Uint16(pkt[36:38]))
})
t.Run("udp checksum of zero goes out as 0xffff", func(t *testing.T) {
pkt := serialize(t, ip4udp, udp(ip4udp), gopacket.Payload{0, 0})
c := binary.BigEndian.Uint16(pkt[26:28])
require.NotZero(t, c)
// The one's complement sum is now 0xffff - c; adding c to the payload
// makes it 0xffff, whose complement is zero.
binary.BigEndian.PutUint16(pkt[28:30], c)
SetTransportChecksum(pkt)
assert.Equal(t, uint16(0xffff), binary.BigEndian.Uint16(pkt[26:28]))
})
}
func TestFold(t *testing.T) {
// 0xffff is the fold's fixed point, so a loop bound one notch tight never
// terminates on it.
for _, tt := range []struct {
in uint32
want uint16
}{
{0, 0},
{0xffff, 0xffff},
{0x10000, 1},
{0x1fffe, 0xffff},
{0xffffffff, 0xffff},
} {
assert.Equal(t, tt.want, fold(tt.in))
}
}
+7
View File
@@ -27,6 +27,13 @@ const (
maxIPv6RejectPacketSize = ipv6.HeaderLen + 8 + 1000
MaxRejectPacketSize = maxIPv6RejectPacketSize
IPProtocolICMP = 1
IPProtocolICMPv6 = 58
IPProtocolTCP = 6
IPProtocolUDP = 17
ICMPv6TypeEchoRequest = 128
ICMPv6TypeEchoReply = 129
)
func CreateRejectPacket(packet []byte, out []byte) []byte {
+15 -1
View File
@@ -34,7 +34,9 @@ type LightHouse struct {
myVpnNetworks []netip.Prefix
myVpnNetworksTable *bart.Lite
punchy *Punchy
// myVpnAddrsTable contains our overlay host addrs, as opposed to the overlay networks
myVpnAddrsTable *bart.Lite
punchy *Punchy
// localAddrsFn enumerates the underlay addresses we advertise. It is a field so tests can supply simulated
// addresses rather than whatever this machine's NICs happen to be. Set it before Start.
@@ -104,6 +106,7 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
amLighthouse: amLighthouse,
myVpnNetworks: cs.myVpnNetworks,
myVpnNetworksTable: cs.myVpnNetworksTable,
myVpnAddrsTable: cs.myVpnAddrsTable,
addrMap: make(map[netip.Addr]*RemoteList),
nebulaPort: nebulaPort,
punchy: p,
@@ -1158,6 +1161,17 @@ func (lhh *LightHouseHandler) handleHostQuery(n *NebulaMeta, fromVpnAddrs []neti
return
}
// Don't respond to requests for us.
if lhh.lh.myVpnAddrsTable.Contains(queryVpnAddr) {
if lhh.l.Enabled(context.Background(), slog.LevelDebug) {
lhh.l.Debug("Ignoring HostQuery for one of my own addresses",
"fromVpnAddrs", fromVpnAddrs,
"queryVpnAddr", queryVpnAddr,
)
}
return
}
found, ln, err := lhh.lh.queryAndPrepMessage(queryVpnAddr, func(c *cache) (int, error) {
n = lhh.resetMeta()
n.Type = NebulaMeta_HostQueryReply
+80 -54
View File
@@ -27,15 +27,27 @@ func TestOldIPv4Only(t *testing.T) {
assert.Equal(t, binary.BigEndian.Uint32(bp[:]), m.GetAddr())
}
func testCertState(networks ...netip.Prefix) *CertState {
cs := &CertState{
myVpnNetworks: networks,
myVpnNetworksTable: new(bart.Lite),
myVpnAddrs: make([]netip.Addr, 0, len(networks)),
myVpnAddrsTable: new(bart.Lite),
}
for _, n := range networks {
cs.myVpnNetworksTable.Insert(n)
cs.myVpnAddrs = append(cs.myVpnAddrs, n.Addr())
cs.myVpnAddrsTable.Insert(netip.PrefixFrom(n.Addr(), n.Addr().BitLen()))
}
return cs
}
func Test_lhStaticMapping(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,
}
cs := testCertState(myVpnNet)
lh1 := "10.128.0.2"
c := config.NewC(l)
@@ -55,12 +67,7 @@ func Test_lhStaticMapping(t *testing.T) {
func TestReloadLighthouseInterval(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,
}
cs := testCertState(myVpnNet)
lh1 := "10.128.0.2"
c := config.NewC(l)
@@ -90,12 +97,7 @@ func TestReloadLighthouseInterval(t *testing.T) {
func BenchmarkLighthouseHandleRequest(b *testing.B) {
l := test.NewLogger()
myVpnNet := netip.MustParsePrefix("10.128.0.1/0")
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
cs := testCertState(myVpnNet)
c := config.NewC(l)
lh, err := NewLightHouseFromConfig(b.Context(), l, c, cs, nil, nil)
@@ -195,12 +197,7 @@ func TestLighthouse_Memory(t *testing.T) {
c.Settings["listen"] = map[string]any{"port": 4242}
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
cs := testCertState(myVpnNet)
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
lh.ifce = &mockEncWriter{}
require.NoError(t, err)
@@ -280,12 +277,7 @@ func TestLighthouse_reload(t *testing.T) {
c.Settings["listen"] = map[string]any{"port": 4242}
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
cs := testCertState(myVpnNet)
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
require.NoError(t, err)
@@ -315,12 +307,7 @@ func TestLighthouse_reloadStaticHostMap(t *testing.T) {
}
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
cs := testCertState(myVpnNet)
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
require.NoError(t, err)
@@ -429,7 +416,9 @@ func TestLighthouse_reloadStaticHostMap(t *testing.T) {
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("3.3.3.3:4242")}, rl.CopyAddrs([]netip.Prefix{}))
}
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
// sendLHHostRequest delivers a HostQuery to lhh and hands back the writer that
// captured what it emitted. Pass a nil filter to see every message.
func sendLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler, filter *NebulaMeta_MessageType) *testEncWriter {
req := &NebulaMeta{
Type: NebulaMeta_HostQuery,
Details: &NebulaMetaDetails{},
@@ -447,12 +436,59 @@ func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, l
panic(err)
}
filter := NebulaMeta_HostQueryReply
w := &testEncWriter{
metaFilter: &filter,
}
w := &testEncWriter{metaFilter: filter}
lhh.HandleRequest(fromAddr, []netip.Addr{myVpnIp}, b, w)
return w.lastReply
return w
}
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
filter := NebulaMeta_HostQueryReply
return sendLHHostRequest(fromAddr, myVpnIp, queryVpnIp, lhh, &filter).lastReply
}
func TestLighthouse_IgnoresHostQueryForItself(t *testing.T) {
// Validate that we don't answer host queries for our own address.
l := test.NewLogger()
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
myVpnIp := myVpnNet.Addr()
c := config.NewC(l)
c.Settings["lighthouse"] = map[string]any{"am_lighthouse": true}
c.Settings["listen"] = map[string]any{"port": 4242}
// Add a static_host_map entry for ourselves, so our address
// is in the addrMap.
c.Settings["static_host_map"] = map[string]any{
myVpnIp.String(): []any{"192.168.100.1:4242"},
}
lh, err := NewLightHouseFromConfig(t.Context(), l, c, testCertState(myVpnNet), nil, nil)
require.NoError(t, err)
lh.ifce = &mockEncWriter{}
lhh := lh.NewRequestHandler()
peerVpnIp := netip.MustParseAddr("10.128.0.2")
peerUdpAddr := netip.MustParseAddrPort("10.0.0.2:4242")
otherVpnIp := netip.MustParseAddr("10.128.0.3")
otherUdpAddr := netip.MustParseAddrPort("10.0.0.3:4242")
newLHHostUpdate(peerUdpAddr, peerVpnIp, []netip.AddrPort{peerUdpAddr}, lhh)
newLHHostUpdate(otherUdpAddr, otherVpnIp, []netip.AddrPort{otherUdpAddr}, lhh)
// Control: a query about a real peer is still answered, and still ends with
// the punch notification aimed at the host that was asked about.
w := sendLHHostRequest(peerUdpAddr, peerVpnIp, otherVpnIp, lhh, nil)
require.NotNil(t, w.lastReply.msg)
assert.Equal(t, NebulaMeta_HostPunchNotification, w.lastReply.msg.Type)
assert.Equal(t, otherVpnIp, w.lastReply.vpnIp)
// Now validate that we don't send to ourselves.
found, _, err := lh.queryAndPrepMessage(myVpnIp, func(*cache) (int, error) { return 0, nil })
require.NoError(t, err)
require.True(t, found, "the lighthouse should hold a cache entry for its own address")
w = sendLHHostRequest(peerUdpAddr, peerVpnIp, myVpnIp, lhh, nil)
assert.Nil(t, w.lastReply.msg, "a query about our own address must produce no reply and no punch notification")
}
func newLHHostUpdate(fromAddr netip.AddrPort, vpnIp netip.Addr, addrs []netip.AddrPort, lhh *LightHouseHandler) {
@@ -642,12 +678,7 @@ func TestLighthouse_Dont_Delete_Static_Hosts(t *testing.T) {
}
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
cs := testCertState(myVpnNet)
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
require.NoError(t, err)
lh.ifce = &mockEncWriter{}
@@ -708,12 +739,7 @@ func TestLighthouse_DeletesWork(t *testing.T) {
}
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
cs := testCertState(myVpnNet)
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
require.NoError(t, err)
lh.ifce = &mockEncWriter{}
+25 -10
View File
@@ -14,6 +14,7 @@ import (
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/cpupick"
"github.com/slackhq/nebula/diag"
"github.com/slackhq/nebula/noiseutil"
"github.com/slackhq/nebula/overlay"
"github.com/slackhq/nebula/sshd"
@@ -68,7 +69,9 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
}
l.Info("Firewall started", "firewallHashes", fw.GetRuleHashes())
ssh, err := sshd.NewSSHServer(ctx, l.With("subsystem", "sshd"))
commands := diag.NewRegistry()
ssh, err := sshd.NewSSHServer(ctx, l.With("subsystem", "sshd"), commands)
if err != nil {
return nil, util.ContextualizeIfNeeded("Error while creating SSH server", err)
}
@@ -244,10 +247,12 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
if pinThreads && routines > 1 && len(cpuAffinity) == 0 && !configTest {
// The operator didn't choose pin CPUs, so pick a default set that
// prefers performance cores and doesn't stack co-located instances
// onto allowed[0]. The bound UDP port keys the per-instance spread:
// distinct across instances sharing a box, stable across restarts.
// A nil result keeps listenIn's stock allowed[i] fallback.
key := uint64(os.Getpid())
// onto allowed[0].
// key is used to seed the spreading of routines->cores.
// use PID if you want to ensure many different Nebulas in VMs or containers land on different cores
// use port if you want to always end up on the same cores, ideal for benchmarking.
key := uint64(os.Getpid()) //default to PID
pinKeyStr := strings.ToLower(c.GetString("tun.pin_threads_key", ""))
switch pinKeyStr {
case "":
@@ -255,14 +260,16 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
case "pid":
l.Debug("tun.pin_threads_key is PID")
case "port":
l.Info("tun.pin_threads_key is port number")
if ap, err := udpConns[0].LocalAddr(); err == nil && ap.Port() != 0 {
l.Info("tun.pin_threads_key is port number")
key = uint64(ap.Port())
} else {
l.Warn("Failed to get a port number for tun.pin_threads_key, falling back to PID", "err", err)
}
default:
l.Warn("tun.pin_threads_key is invalid, using PID")
}
if ap, err := udpConns[0].LocalAddr(); err == nil && ap.Port() != 0 {
key = uint64(ap.Port())
}
cpuAffinity = cpupick.Default(routines, key, l)
}
@@ -318,13 +325,20 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
return nil, util.ContextualizeIfNeeded("Failed to start stats emitter", err)
}
// Built before the configTest return so that a bad ctl block fails `nebula -test`. It only
// holds the registry, which attachCommands populates below, and reads nothing until Start.
ctlServer, err := newCtlServerFromConfig(ctx, l.With("subsystem", "ctl"), c, commands)
if err != nil {
return nil, util.ContextualizeIfNeeded("Failed to configure the ctl socket", err)
}
if configTest {
return nil, nil
}
go ifce.emitStats(ctx, c.GetDuration("stats.interval", time.Second*10))
attachCommands(l, c, ssh, ifce)
attachCommands(l, c, commands, ifce)
networkChanges := udp.NewNetworkChangeMonitor(ctx, l, c)
@@ -335,6 +349,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
ctx: ctx,
cancel: cancel,
sshStart: sshStart,
ctlStart: ctlServer.Start,
statsStart: stats.Start,
dnsStart: ds.Start,
lighthouseStart: lightHouse.StartUpdateWorker,
+6 -7
View File
@@ -8,7 +8,6 @@ import (
"net/netip"
"time"
"github.com/google/gopacket/layers"
"golang.org/x/net/ipv6"
"github.com/slackhq/nebula/firewall"
@@ -369,15 +368,15 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
return nil
}
switch layers.IPProtocol(proto) {
case layers.IPProtocolICMPv6:
switch proto {
case iputil.IPProtocolICMPv6:
// An ICMPv6 message is at least type, code and checksum, 4 bytes. Only echo carries more than we read.
if dataLen < offset+4 {
return ErrIPv6PacketTooShort
}
fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6
switch data[offset] { //icmp type
case layers.ICMPv6TypeEchoRequest, layers.ICMPv6TypeEchoReply:
case iputil.ICMPv6TypeEchoRequest, iputil.ICMPv6TypeEchoReply:
if dataLen < offset+6 {
return ErrIPv6PacketTooShort
}
@@ -386,7 +385,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
fp.RemotePort = 0
}
case layers.IPProtocolTCP, layers.IPProtocolUDP:
case iputil.IPProtocolTCP, iputil.IPProtocolUDP:
if dataLen < offset+4 {
return ErrIPv6PacketTooShort
}
@@ -435,7 +434,7 @@ func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
// Accounting for a variable header length, do we have enough data for our src/dst tuples?
minLen := ihl
if !fp.Fragment {
if fp.Protocol == firewall.ProtoICMP {
if fp.Protocol == iputil.IPProtocolICMP {
minLen += minFwPacketLen + 2
} else {
minLen += minFwPacketLen
@@ -457,7 +456,7 @@ func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
if fp.Fragment {
fp.RemotePort = 0
fp.LocalPort = 0
} else if fp.Protocol == firewall.ProtoICMP { //note that orientation doesn't matter on ICMP
} else if fp.Protocol == iputil.IPProtocolICMP { //note that orientation doesn't matter on ICMP
fp.RemotePort = binary.BigEndian.Uint16(data[ihl+4 : ihl+6]) //identifier
fp.LocalPort = 0 //code would be uint16(data[ihl+1])
} else if incoming {
+20 -19
View File
@@ -9,6 +9,7 @@ import (
"github.com/google/gopacket"
"github.com/google/gopacket/layers"
"github.com/slackhq/nebula/iputil"
"github.com/slackhq/nebula/firewall"
"github.com/stretchr/testify/assert"
@@ -58,7 +59,7 @@ func Test_newPacket(t *testing.T) {
Src: net.IPv4(10, 0, 0, 1),
Dst: net.IPv4(10, 0, 0, 2),
Options: []byte{0, 1, 0, 2},
Protocol: firewall.ProtoTCP,
Protocol: iputil.IPProtocolTCP,
}
b, _ = h.Marshal()
@@ -66,7 +67,7 @@ func Test_newPacket(t *testing.T) {
err = newPacket(b, true, p)
require.NoError(t, err)
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
assert.Equal(t, netip.MustParseAddr("10.0.0.2"), p.LocalAddr)
assert.Equal(t, netip.MustParseAddr("10.0.0.1"), p.RemoteAddr)
assert.Equal(t, uint16(3), p.RemotePort)
@@ -239,7 +240,7 @@ func Test_newPacket_v6(t *testing.T) {
// A good UDP packet
ip = layers.IPv6{
Version: 6,
NextHeader: firewall.ProtoUDP,
NextHeader: iputil.IPProtocolUDP,
HopLimit: 128,
SrcIP: net.IPv6linklocalallrouters,
DstIP: net.IPv6linklocalallnodes,
@@ -262,7 +263,7 @@ func Test_newPacket_v6(t *testing.T) {
// incoming
err = newPacket(b, true, p)
require.NoError(t, err)
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
assert.Equal(t, uint16(36123), p.RemotePort)
@@ -272,7 +273,7 @@ func Test_newPacket_v6(t *testing.T) {
// outgoing
err = newPacket(b, false, p)
require.NoError(t, err)
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.RemoteAddr)
assert.Equal(t, uint16(36123), p.LocalPort)
@@ -289,7 +290,7 @@ func Test_newPacket_v6(t *testing.T) {
// incoming
err = newPacket(b, true, p)
require.NoError(t, err)
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
assert.Equal(t, uint16(36123), p.RemotePort)
@@ -299,7 +300,7 @@ func Test_newPacket_v6(t *testing.T) {
// outgoing
err = newPacket(b, false, p)
require.NoError(t, err)
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.RemoteAddr)
assert.Equal(t, uint16(36123), p.LocalPort)
@@ -344,7 +345,7 @@ func Test_newPacket_v6(t *testing.T) {
err = newPacket(b, true, p)
require.NoError(t, err)
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
assert.Equal(t, uint16(36123), p.RemotePort)
@@ -678,7 +679,7 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
pkt := make([]byte, realTCPAt+4)
pkt[0] = 0x60 // version 6
pkt[6] = byte(layers.IPProtocolIPv6Destination) // NextHeader -> Destination Options
pkt[40] = byte(firewall.ProtoTCP) // Dest-Options NextHeader -> TCP
pkt[40] = byte(iputil.IPProtocolTCP) // Dest-Options NextHeader -> TCP
pkt[41] = 255 // HdrExtLen = 255
// Forged transport header at the pre-fix (wrong) offset: dst port 443.
@@ -687,7 +688,7 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
binary.BigEndian.PutUint16(pkt[realTCPAt+2:realTCPAt+4], 22)
require.NoError(t, newPacket(pkt, true, p))
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
// LocalPort is the destination port for incoming traffic. It must be the real port (22)
// the host delivers to, not the forged 443 at the overflowed offset.
assert.Equal(t, uint16(22), p.LocalPort, "firewall must parse the real transport header, not the overflowed offset")
@@ -766,7 +767,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
// Plain IPv4 TCP, IHL 20: L4 offset 20, no fragment shape.
v4 := make([]byte, 28)
v4[0] = 0x45
v4[9] = firewall.ProtoTCP
v4[9] = iputil.IPProtocolTCP
binary.BigEndian.PutUint16(v4[6:8], 0x4000) // DF only
require.NoError(t, newPacket(v4, true, p))
assert.Equal(t, 20, p.IPHdrLen)
@@ -777,7 +778,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
// (Fragment false) but the coalescer must not touch it (FragAny true).
ff := make([]byte, 28)
ff[0] = 0x45
ff[9] = firewall.ProtoUDP
ff[9] = iputil.IPProtocolUDP
binary.BigEndian.PutUint16(ff[6:8], 0x2000) // MF, offset 0
require.NoError(t, newPacket(ff, true, p))
assert.False(t, p.Fragment)
@@ -787,7 +788,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
// IPv4 non-first fragment (nonzero offset): both flags set.
nf := make([]byte, 28)
nf[0] = 0x45
nf[9] = firewall.ProtoUDP
nf[9] = iputil.IPProtocolUDP
binary.BigEndian.PutUint16(nf[6:8], 0x00b9)
require.NoError(t, newPacket(nf, true, p))
assert.True(t, p.Fragment)
@@ -796,7 +797,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
// IPv4 with options (IHL 24): IPHdrLen tracks the real L4 offset.
opts := make([]byte, 32)
opts[0] = 0x46
opts[9] = firewall.ProtoTCP
opts[9] = iputil.IPProtocolTCP
binary.BigEndian.PutUint16(opts[6:8], 0x4000)
require.NoError(t, newPacket(opts, true, p))
assert.Equal(t, 24, p.IPHdrLen)
@@ -805,7 +806,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
// Plain IPv6 TCP: L4 at 40.
v6 := make([]byte, 60)
v6[0] = 0x60
v6[6] = firewall.ProtoTCP
v6[6] = iputil.IPProtocolTCP
require.NoError(t, newPacket(v6, true, p))
assert.Equal(t, 40, p.IPHdrLen)
assert.False(t, p.FragAny)
@@ -814,7 +815,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
hbh := make([]byte, 60)
hbh[0] = 0x60
hbh[6] = 0 // hop-by-hop
hbh[40] = firewall.ProtoTCP
hbh[40] = iputil.IPProtocolTCP
hbh[41] = 0 // HdrExtLen 0 -> 8-byte header
require.NoError(t, newPacket(hbh, true, p))
assert.Equal(t, 48, p.IPHdrLen)
@@ -824,17 +825,17 @@ func Test_newPacket_parsedFields(t *testing.T) {
f6 := make([]byte, 60)
f6[0] = 0x60
f6[6] = 44 // fragment extension header
f6[40] = firewall.ProtoUDP
f6[40] = iputil.IPProtocolUDP
require.NoError(t, newPacket(f6, true, p))
assert.True(t, p.FragAny)
assert.False(t, p.Fragment)
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
// IPv6 non-first fragment: both set, walk stops at the fragment header.
f6n := make([]byte, 60)
f6n[0] = 0x60
f6n[6] = 44
f6n[40] = firewall.ProtoUDP
f6n[40] = iputil.IPProtocolUDP
binary.BigEndian.PutUint16(f6n[42:44], 0x0008)
require.NoError(t, newPacket(f6n, true, p))
assert.True(t, p.Fragment)
+29
View File
@@ -11,6 +11,7 @@ import (
"os"
"path/filepath"
"runtime"
"slices"
"sync/atomic"
"syscall"
"unsafe"
@@ -182,6 +183,7 @@ func (t *winTun) addRoutes(logErrors bool) error {
luid := winipcfg.LUID(t.tun.LUID())
routes := *t.Routes.Load()
foundDefault4 := false
carriesV6 := slices.ContainsFunc(t.vpnNetworks, func(p netip.Prefix) bool { return p.Addr().Is6() })
for _, r := range routes {
if len(r.Via) == 0 || !r.Install {
@@ -189,6 +191,9 @@ func (t *winTun) addRoutes(logErrors bool) error {
continue
}
// A v6 unsafe_route is legal under a v4-only cert; uninstalled ones put nothing on the adapter.
carriesV6 = carriesV6 || r.Cidr.Addr().Is6()
// Add our unsafe route as an on-link route to the nebula tun device.
err := luid.AddRoute(r.Cidr, unspecifiedNextHop(r.Cidr), uint32(r.Metric))
if err != nil {
@@ -210,6 +215,11 @@ func (t *winTun) addRoutes(logErrors bool) error {
}
}
return t.setMTU(luid, foundDefault4, carriesV6)
}
// setMTU applies tun.mtu per address family. The default route metric rides along on the v4 handle.
func (t *winTun) setMTU(luid winipcfg.LUID, foundDefault4, carriesV6 bool) error {
ipif, err := luid.IPInterface(windows.AF_INET)
if err != nil {
return fmt.Errorf("failed to get ip interface: %w", err)
@@ -224,6 +234,25 @@ func (t *winTun) addRoutes(logErrors bool) error {
if err := ipif.Set(); err != nil {
return fmt.Errorf("failed to set ip interface: %w", err)
}
// Windows tracks NLMTU per family and wintun sets neither, so v6 keeps the adapter default of 65535.
// Gated so a v4-only overlay under 1280 boots; a v6 one deliberately does not, as linux also refuses.
if !carriesV6 {
return nil
}
ipif6, err := luid.IPInterface(windows.AF_INET6)
if err != nil {
// No v6 on the adapter means there is no NLMTU to get wrong. A failed Set below is not the same thing.
t.l.Info("Skipping ipv6 MTU, no ipv6 interface on this adapter", "error", err)
return nil
}
ipif6.NLMTU = uint32(t.MTU)
if err := ipif6.Set(); err != nil {
return fmt.Errorf("failed to set ipv6 interface: %w", err)
}
return nil
}
+3 -905
View File
@@ -1,62 +1,19 @@
package nebula
// Configuration and lifecycle for the ssh debug console. The commands it serves are not
// defined here; see commands.go, which registers them for every transport.
import (
"bytes"
"encoding/json"
"errors"
"flag"
"fmt"
"log/slog"
"maps"
"net"
"net/netip"
"os"
"path/filepath"
"runtime"
"runtime/pprof"
"sort"
"strconv"
"strings"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/logging"
"github.com/slackhq/nebula/sshd"
)
type sshListHostMapFlags struct {
Json bool
Pretty bool
ByIndex bool
}
type sshPrintCertFlags struct {
Json bool
Pretty bool
Raw bool
}
type sshPrintTunnelFlags struct {
Pretty bool
}
type sshChangeRemoteFlags struct {
Address string
}
type sshCloseTunnelFlags struct {
LocalOnly bool
}
type sshCreateTunnelFlags struct {
Address string
}
type sshDeviceInfoFlags struct {
Json bool
Pretty bool
}
func wireSSHReload(l *slog.Logger, ssh *sshd.SSHServer, c *config.C) {
c.RegisterReloadCallback(func(c *config.C) {
if c.GetBool("sshd.enabled", false) {
@@ -197,862 +154,3 @@ func configSSH(l *slog.Logger, ssh *sshd.SSHServer, c *config.C) (func(), error)
return runner, nil
}
func attachCommands(l *slog.Logger, c *config.C, ssh *sshd.SSHServer, f *Interface) {
// sandboxDir defaults to a dir in temp. The intention is that end user will
// create this dir as needed. Overriding this config value to "" allows
// writing to anywhere in the system.
defaultDir := filepath.Join(os.TempDir(), "nebula-debug")
sandboxDir := c.GetString("sshd.sandbox_dir", defaultDir)
ssh.RegisterCommand(&sshd.Command{
Name: "list-hostmap",
ShortDescription: "List all known previously connected hosts",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := sshListHostMapFlags{}
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
fl.BoolVar(&s.ByIndex, "by-index", false, "gets all hosts in the hostmap from the index table")
return fl, &s
},
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshListHostMap(f.hostMap, fs, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "list-pending-hostmap",
ShortDescription: "List all handshaking hosts",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := sshListHostMapFlags{}
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
fl.BoolVar(&s.ByIndex, "by-index", false, "gets all hosts in the hostmap from the index table")
return fl, &s
},
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshListHostMap(f.handshakeManager, fs, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "list-lighthouse-addrmap",
ShortDescription: "List all lighthouse map entries",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := sshListHostMapFlags{}
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
return fl, &s
},
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshListLighthouseMap(f.lightHouse, fs, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "reload",
ShortDescription: "Reloads configuration from disk, same as sending HUP to the process",
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshReload(c, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "start-cpu-profile",
ShortDescription: "Starts a cpu profile and write output to the provided file, ex: `cpu-profile.pb.gz`",
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshStartCpuProfile(sandboxDir, fs, a, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "stop-cpu-profile",
ShortDescription: "Stops a cpu profile and writes output to the previously provided file",
Callback: func(fs any, a []string, w sshd.StringWriter) error {
pprof.StopCPUProfile()
return w.WriteLine("If a CPU profile was running it is now stopped")
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "save-heap-profile",
ShortDescription: "Saves a heap profile to the provided path, ex: `heap-profile.pb.gz`",
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshGetHeapProfile(sandboxDir, fs, a, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "mutex-profile-fraction",
ShortDescription: "Gets or sets runtime.SetMutexProfileFraction",
Callback: sshMutexProfileFraction,
})
ssh.RegisterCommand(&sshd.Command{
Name: "save-mutex-profile",
ShortDescription: "Saves a mutex profile to the provided path, ex: `mutex-profile.pb.gz`",
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshGetMutexProfile(sandboxDir, fs, a, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "log-level",
ShortDescription: "Gets or sets the current log level",
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshLogLevel(l, fs, a, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "log-format",
ShortDescription: "Gets or sets the current log format",
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshLogFormat(l, fs, a, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "version",
ShortDescription: "Prints the currently running version of nebula",
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshVersion(f, fs, a, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "device-info",
ShortDescription: "Prints information about the network device.",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := sshDeviceInfoFlags{}
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
return fl, &s
},
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshDeviceInfo(f, fs, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "print-cert",
ShortDescription: "Prints the current certificate being used or the certificate for the provided vpn addr",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := sshPrintCertFlags{}
fl.BoolVar(&s.Json, "json", false, "outputs as json")
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
fl.BoolVar(&s.Raw, "raw", false, "raw prints the PEM encoded certificate, not compatible with -json or -pretty")
return fl, &s
},
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshPrintCert(f, fs, a, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "print-tunnel",
ShortDescription: "Prints json details about a tunnel for the provided vpn addr",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := sshPrintTunnelFlags{}
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json")
return fl, &s
},
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshPrintTunnel(f, fs, a, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "print-relays",
ShortDescription: "Prints json details about all relay info",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := sshPrintTunnelFlags{}
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json")
return fl, &s
},
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshPrintRelays(f, fs, a, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "change-remote",
ShortDescription: "Changes the remote address used in the tunnel for the provided vpn addr",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := sshChangeRemoteFlags{}
fl.StringVar(&s.Address, "address", "", "The new remote address, ip:port")
return fl, &s
},
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshChangeRemote(f, fs, a, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "close-tunnel",
ShortDescription: "Closes a tunnel for the provided vpn addr",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := sshCloseTunnelFlags{}
fl.BoolVar(&s.LocalOnly, "local-only", false, "Disables notifying the remote that the tunnel is shutting down")
return fl, &s
},
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshCloseTunnel(f, fs, a, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "create-tunnel",
ShortDescription: "Creates a tunnel for the provided vpn address",
Help: "The lighthouses will be queried for real addresses but you can provide one as well.",
Flags: func() (*flag.FlagSet, any) {
fl := flag.NewFlagSet("", flag.ContinueOnError)
s := sshCreateTunnelFlags{}
fl.StringVar(&s.Address, "address", "", "Optionally provide a real remote address, ip:port ")
return fl, &s
},
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshCreateTunnel(f, fs, a, w)
},
})
ssh.RegisterCommand(&sshd.Command{
Name: "query-lighthouse",
ShortDescription: "Query the lighthouses for the provided vpn address",
Help: "This command is asynchronous. Only currently known udp addresses will be printed.",
Callback: func(fs any, a []string, w sshd.StringWriter) error {
return sshQueryLighthouse(f, fs, a, w)
},
})
}
func sshListHostMap(hl controlHostLister, a any, w sshd.StringWriter) error {
fs, ok := a.(*sshListHostMapFlags)
if !ok {
return nil
}
var hm []ControlHostInfo
if fs.ByIndex {
hm = listHostMapIndexes(hl)
} else {
hm = listHostMapHosts(hl)
}
sort.Slice(hm, func(i, j int) bool {
return hm[i].VpnAddrs[0].Compare(hm[j].VpnAddrs[0]) < 0
})
if fs.Json || fs.Pretty {
js := json.NewEncoder(w.GetWriter())
if fs.Pretty {
js.SetIndent("", " ")
}
err := js.Encode(hm)
if err != nil {
return nil
}
} else {
for _, v := range hm {
err := w.WriteLine(fmt.Sprintf("%s: %s", v.VpnAddrs, v.RemoteAddrs))
if err != nil {
return err
}
}
}
return nil
}
func sshListLighthouseMap(lightHouse *LightHouse, a any, w sshd.StringWriter) error {
fs, ok := a.(*sshListHostMapFlags)
if !ok {
return nil
}
type lighthouseInfo struct {
VpnAddr string `json:"vpnAddr"`
Addrs *CacheMap `json:"addrs"`
}
lightHouse.RLock()
addrMap := make([]lighthouseInfo, len(lightHouse.addrMap))
x := 0
for k, v := range lightHouse.addrMap {
addrMap[x] = lighthouseInfo{
VpnAddr: k.String(),
Addrs: v.CopyCache(),
}
x++
}
lightHouse.RUnlock()
sort.Slice(addrMap, func(i, j int) bool {
return strings.Compare(addrMap[i].VpnAddr, addrMap[j].VpnAddr) < 0
})
if fs.Json || fs.Pretty {
js := json.NewEncoder(w.GetWriter())
if fs.Pretty {
js.SetIndent("", " ")
}
err := js.Encode(addrMap)
if err != nil {
return nil
}
} else {
for _, v := range addrMap {
b, err := json.Marshal(v.Addrs)
if err != nil {
return err
}
err = w.WriteLine(fmt.Sprintf("%s: %s", v.VpnAddr, string(b)))
if err != nil {
return err
}
}
}
return nil
}
// sshSanitizeFilePath validates that the given file path is within the sandbox directory.
// If sandboxDir is empty, the path is returned as-is for backwards compatibility.
func sshSanitizeFilePath(sandboxDir, filePath string) (string, error) {
if sandboxDir == "" {
return filePath, nil
}
// Clean and resolve the path relative to the sandbox directory
if !filepath.IsAbs(filePath) {
filePath = filepath.Join(sandboxDir, filePath)
}
cleaned := filepath.Clean(filePath)
// Ensure the resolved path is within the sandbox directory
cleanedSandbox := filepath.Clean(sandboxDir)
if cleaned == cleanedSandbox {
return "", fmt.Errorf("path %q resolves to the sandbox directory itself %q", filePath, sandboxDir)
}
if !strings.HasPrefix(cleaned, cleanedSandbox+string(filepath.Separator)) {
return "", fmt.Errorf("path %q is outside the sandbox directory %q", filePath, sandboxDir)
}
return cleaned, nil
}
func sshStartCpuProfile(sandboxDir string, fs any, a []string, w sshd.StringWriter) error {
if len(a) == 0 {
err := w.WriteLine("No path to write profile provided")
return err
}
filePath, err := sshSanitizeFilePath(sandboxDir, a[0])
if err != nil {
return w.WriteLine(err.Error())
}
file, err := os.Create(filePath)
if err != nil {
err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
return err
}
err = pprof.StartCPUProfile(file)
if err != nil {
err = w.WriteLine(fmt.Sprintf("Unable to start cpu profile: %s", err))
return err
}
err = w.WriteLine(fmt.Sprintf("Started cpu profile, issue stop-cpu-profile to write the output to %s", a))
return err
}
func sshVersion(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
return w.WriteLine(fmt.Sprintf("%s", ifce.version))
}
func sshQueryLighthouse(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
if len(a) == 0 {
return w.WriteLine("No vpn address was provided")
}
vpnAddr, err := netip.ParseAddr(a[0])
if err != nil {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
if !vpnAddr.IsValid() {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
var cm *CacheMap
rl := ifce.lightHouse.Query(vpnAddr)
if rl != nil {
cm = rl.CopyCache()
}
return json.NewEncoder(w.GetWriter()).Encode(cm)
}
func sshCloseTunnel(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
flags, ok := fs.(*sshCloseTunnelFlags)
if !ok {
return nil
}
if len(a) == 0 {
return w.WriteLine("No vpn address was provided")
}
vpnAddr, err := netip.ParseAddr(a[0])
if err != nil {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
if !vpnAddr.IsValid() {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
if hostInfo == nil {
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn address: %v", a[0]))
}
if !flags.LocalOnly {
ifce.send(
header.CloseTunnel,
0,
hostInfo.ConnectionState,
hostInfo,
[]byte{},
make([]byte, 12, 12),
make([]byte, mtu),
)
}
ifce.closeTunnel(hostInfo)
return w.WriteLine("Closed")
}
func sshCreateTunnel(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
flags, ok := fs.(*sshCreateTunnelFlags)
if !ok {
return nil
}
if len(a) == 0 {
return w.WriteLine("No vpn address was provided")
}
vpnAddr, err := netip.ParseAddr(a[0])
if err != nil {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
if !vpnAddr.IsValid() {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
if hostInfo != nil {
return w.WriteLine(fmt.Sprintf("Tunnel already exists"))
}
hostInfo = ifce.handshakeManager.QueryVpnAddr(vpnAddr)
if hostInfo != nil {
return w.WriteLine(fmt.Sprintf("Tunnel already handshaking"))
}
var addr netip.AddrPort
if flags.Address != "" {
addr, err = netip.ParseAddrPort(flags.Address)
if err != nil {
return w.WriteLine("Address could not be parsed")
}
}
hostInfo = ifce.handshakeManager.StartHandshake(vpnAddr, nil)
if addr.IsValid() {
hostInfo.SetRemote(addr)
}
return w.WriteLine("Created")
}
func sshChangeRemote(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
flags, ok := fs.(*sshChangeRemoteFlags)
if !ok {
return nil
}
if len(a) == 0 {
return w.WriteLine("No vpn address was provided")
}
if flags.Address == "" {
return w.WriteLine("No address was provided")
}
addr, err := netip.ParseAddrPort(flags.Address)
if err != nil {
return w.WriteLine("Address could not be parsed")
}
vpnAddr, err := netip.ParseAddr(a[0])
if err != nil {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
if !vpnAddr.IsValid() {
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
}
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
if hostInfo == nil {
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn address: %v", a[0]))
}
hostInfo.SetRemote(addr)
return w.WriteLine("Changed")
}
func sshGetHeapProfile(sandboxDir string, fs any, a []string, w sshd.StringWriter) error {
if len(a) == 0 {
return w.WriteLine("No path to write profile provided")
}
filePath, err := sshSanitizeFilePath(sandboxDir, a[0])
if err != nil {
return w.WriteLine(err.Error())
}
file, err := os.Create(filePath)
if err != nil {
err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
return err
}
err = pprof.WriteHeapProfile(file)
if err != nil {
err = w.WriteLine(fmt.Sprintf("Unable to write profile: %s", err))
return err
}
err = w.WriteLine(fmt.Sprintf("Mem profile created at %s", a))
return err
}
func sshMutexProfileFraction(fs any, a []string, w sshd.StringWriter) error {
if len(a) == 0 {
rate := runtime.SetMutexProfileFraction(-1)
return w.WriteLine(fmt.Sprintf("Current value: %d", rate))
}
newRate, err := strconv.Atoi(a[0])
if err != nil {
return w.WriteLine(fmt.Sprintf("Invalid argument: %s", a[0]))
}
oldRate := runtime.SetMutexProfileFraction(newRate)
return w.WriteLine(fmt.Sprintf("New value: %d. Old value: %d", newRate, oldRate))
}
func sshGetMutexProfile(sandboxDir string, fs any, a []string, w sshd.StringWriter) error {
if len(a) == 0 {
return w.WriteLine("No path to write profile provided")
}
filePath, err := sshSanitizeFilePath(sandboxDir, a[0])
if err != nil {
return w.WriteLine(err.Error())
}
file, err := os.Create(filePath)
if err != nil {
return w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
}
defer file.Close()
mutexProfile := pprof.Lookup("mutex")
if mutexProfile == nil {
return w.WriteLine("Unable to get pprof.Lookup(\"mutex\")")
}
err = mutexProfile.WriteTo(file, 0)
if err != nil {
return w.WriteLine(fmt.Sprintf("Unable to write profile: %s", err))
}
return w.WriteLine(fmt.Sprintf("Mutex profile created at %s", a))
}
func sshLogLevel(l *slog.Logger, fs any, a []string, w sshd.StringWriter) error {
ctrl, ok := l.Handler().(interface {
GetLevel() slog.Level
SetLevel(slog.Level)
})
if !ok {
return w.WriteLine("Log level is not reconfigurable on this logger")
}
if len(a) == 0 {
return w.WriteLine(fmt.Sprintf("Log level is: %s", logging.LevelName(ctrl.GetLevel())))
}
level, err := logging.ParseLevel(strings.ToLower(a[0]))
if err != nil {
return w.WriteLine(fmt.Sprintf("Unknown log level %s. Possible log levels: trace, debug, info, warn, error", a))
}
ctrl.SetLevel(level)
return w.WriteLine(fmt.Sprintf("Log level is: %s", logging.LevelName(ctrl.GetLevel())))
}
func sshLogFormat(l *slog.Logger, fs any, a []string, w sshd.StringWriter) error {
ctrl, ok := l.Handler().(interface {
GetFormat() string
SetFormat(string) error
})
if !ok {
return w.WriteLine("Log format is not reconfigurable on this logger")
}
if len(a) == 0 {
return w.WriteLine(fmt.Sprintf("Log format is: %s", ctrl.GetFormat()))
}
if err := ctrl.SetFormat(strings.ToLower(a[0])); err != nil {
return err
}
return w.WriteLine(fmt.Sprintf("Log format is: %s", ctrl.GetFormat()))
}
func sshPrintCert(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
args, ok := fs.(*sshPrintCertFlags)
if !ok {
return nil
}
cert := ifce.pki.getCertState().GetDefaultCertificate()
if len(a) > 0 {
vpnAddr, err := netip.ParseAddr(a[0])
if err != nil {
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
}
if !vpnAddr.IsValid() {
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
}
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
if hostInfo == nil {
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn addr: %v", a[0]))
}
cert = hostInfo.GetCert().Certificate
}
if args.Json || args.Pretty {
b, err := cert.MarshalJSON()
if err != nil {
return nil
}
if args.Pretty {
buf := new(bytes.Buffer)
err := json.Indent(buf, b, "", " ")
b = buf.Bytes()
if err != nil {
return nil
}
}
return w.WriteBytes(b)
}
if args.Raw {
b, err := cert.MarshalPEM()
if err != nil {
return nil
}
return w.WriteBytes(b)
}
return w.WriteLine(cert.String())
}
func sshPrintRelays(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
args, ok := fs.(*sshPrintTunnelFlags)
if !ok {
w.WriteLine(fmt.Sprintf("sshPrintRelays failed to convert args type"))
return nil
}
relays := map[uint32]*HostInfo{}
ifce.hostMap.Lock()
maps.Copy(relays, ifce.hostMap.Relays)
ifce.hostMap.Unlock()
type RelayFor struct {
Error error
Type string
State string
PeerAddr netip.Addr
LocalIndex uint32
RemoteIndex uint32
RelayedThrough []netip.Addr
}
type RelayOutput struct {
NebulaAddr netip.Addr
RelayForAddrs []RelayFor
}
type CmdOutput struct {
Relays []*RelayOutput
}
co := CmdOutput{}
enc := json.NewEncoder(w.GetWriter())
if args.Pretty {
enc.SetIndent("", " ")
}
for k, v := range relays {
ro := RelayOutput{NebulaAddr: v.vpnAddrs[0]}
co.Relays = append(co.Relays, &ro)
relayHI := ifce.hostMap.QueryVpnAddr(v.vpnAddrs[0])
if relayHI == nil {
ro.RelayForAddrs = append(ro.RelayForAddrs, RelayFor{Error: errors.New("could not find hostinfo")})
continue
}
for _, vpnAddr := range relayHI.relayState.CopyRelayForIps() {
rf := RelayFor{Error: nil}
r, ok := relayHI.relayState.GetRelayForByAddr(vpnAddr)
if ok {
t := ""
switch r.Type {
case ForwardingType:
t = "forwarding"
case TerminalType:
t = "terminal"
default:
t = "unknown"
}
s := ""
switch r.State {
case Requested:
s = "requested"
case Established:
s = "established"
default:
s = "unknown"
}
rf.LocalIndex = r.LocalIndex
rf.RemoteIndex = r.RemoteIndex
rf.PeerAddr = r.PeerAddr
rf.Type = t
rf.State = s
if rf.LocalIndex != k {
rf.Error = fmt.Errorf("hostmap LocalIndex '%v' does not match RelayState LocalIndex", k)
}
}
relayedHI := ifce.hostMap.QueryVpnAddr(vpnAddr)
if relayedHI != nil {
rf.RelayedThrough = append(rf.RelayedThrough, relayedHI.relayState.CopyRelayIps()...)
}
ro.RelayForAddrs = append(ro.RelayForAddrs, rf)
}
}
err := enc.Encode(co)
if err != nil {
return err
}
return nil
}
func sshPrintTunnel(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
args, ok := fs.(*sshPrintTunnelFlags)
if !ok {
return nil
}
if len(a) == 0 {
return w.WriteLine("No vpn address was provided")
}
vpnAddr, err := netip.ParseAddr(a[0])
if err != nil {
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
}
if !vpnAddr.IsValid() {
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
}
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
if hostInfo == nil {
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn addr: %v", a[0]))
}
enc := json.NewEncoder(w.GetWriter())
if args.Pretty {
enc.SetIndent("", " ")
}
return enc.Encode(copyHostInfo(hostInfo, ifce.hostMap.GetPreferredRanges()))
}
func sshDeviceInfo(ifce *Interface, fs any, w sshd.StringWriter) error {
data := struct {
Name string `json:"name"`
Cidr []netip.Prefix `json:"cidr"`
}{
Name: ifce.inside.Name(),
Cidr: make([]netip.Prefix, len(ifce.inside.Networks())),
}
copy(data.Cidr, ifce.inside.Networks())
flags, ok := fs.(*sshDeviceInfoFlags)
if !ok {
return fmt.Errorf("internal error: expected flags to be sshDeviceInfoFlags but was %+v", fs)
}
if flags.Json || flags.Pretty {
js := json.NewEncoder(w.GetWriter())
if flags.Pretty {
js.SetIndent("", " ")
}
return js.Encode(data)
} else {
return w.WriteLine(fmt.Sprintf("name=%v cidr=%v", data.Name, data.Cidr))
}
}
func sshReload(c *config.C, w sshd.StringWriter) error {
err := w.WriteLine("Reloading config")
c.ReloadConfig()
return err
}
+7 -20
View File
@@ -9,8 +9,9 @@ import (
"net"
"sync"
"github.com/armon/go-radix"
"golang.org/x/crypto/ssh"
"github.com/slackhq/nebula/diag"
)
type SSHServer struct {
@@ -25,10 +26,9 @@ type SSHServer struct {
trustedKeys map[string]map[string]bool
trustedCAs []ssh.PublicKey
// List of available commands
helpCommand *Command
commands *radix.Tree
listener net.Listener
// The commands this server serves. Shared with every other transport, see diag.Registry.
commands *diag.Registry
listener net.Listener
// ctx parents per-Run contexts. Cancelling it (e.g. via Control.Stop) tears the server down even
// across reloads, since each Run derives a fresh child rather than reusing this one directly.
@@ -38,11 +38,11 @@ type SSHServer struct {
// NewSSHServer creates a new ssh server rigged with default commands and prepares to listen.
// The ssh server's context is parented off the supplied ctx so cancelling it
// (e.g. on Control.Stop) tears down active sessions and closes the listener.
func NewSSHServer(ctx context.Context, l *slog.Logger) (*SSHServer, error) {
func NewSSHServer(ctx context.Context, l *slog.Logger, commands *diag.Registry) (*SSHServer, error) {
s := &SSHServer{
trustedKeys: make(map[string]map[string]bool),
l: l,
commands: radix.New(),
commands: commands,
ctx: ctx,
}
@@ -90,14 +90,6 @@ func NewSSHServer(ctx context.Context, l *slog.Logger) (*SSHServer, error) {
ServerVersion: fmt.Sprintf("SSH-2.0-Nebula???"),
}
s.RegisterCommand(&Command{
Name: "help",
ShortDescription: "prints available commands or help <command> for specific usage info",
Callback: func(a any, args []string, w StringWriter) error {
return helpCallback(s.commands, args, w)
},
})
return s, nil
}
@@ -160,11 +152,6 @@ func (s *SSHServer) AddAuthorizedKey(user, pubKey string) error {
return nil
}
// RegisterCommand adds a command that can be run by a user, by default only `help` is available
func (s *SSHServer) RegisterCommand(c *Command) {
s.commands.Insert(c.Name, c)
}
// Run begins listening and accepting connections. Each invocation derives a fresh per-Run context
// from the constructor-supplied ctx so a Stop+Run sequence (used by config reload) starts clean
// rather than carrying a permanently-cancelled context across runs.
+18 -45
View File
@@ -1,37 +1,38 @@
package sshd
import (
"fmt"
"log/slog"
"sort"
"strings"
"github.com/anmitsu/go-shlex"
"github.com/armon/go-radix"
"golang.org/x/crypto/ssh"
"golang.org/x/term"
"github.com/slackhq/nebula/diag"
)
type session struct {
l *slog.Logger
c *ssh.ServerConn
term *term.Terminal
commands *radix.Tree
commands *diag.Registry
cancel func()
}
func NewSession(commands *radix.Tree, conn *ssh.ServerConn, chans <-chan ssh.NewChannel, cancel func(), l *slog.Logger) *session {
func NewSession(commands *diag.Registry, conn *ssh.ServerConn, chans <-chan ssh.NewChannel, cancel func(), l *slog.Logger) *session {
s := &session{
commands: radix.NewFromMap(commands.ToMap()),
// A copy, so the logout command this session adds for itself stays invisible to every
// other session and to `nebula ctl`.
commands: commands.Clone(),
l: l,
c: conn,
cancel: cancel,
}
s.commands.Insert("logout", &Command{
s.commands.RegisterCommand(&diag.Command{
Name: "logout",
ShortDescription: "Ends the current session",
Callback: func(a any, args []string, w StringWriter) error {
Callback: func(a any, args []string, w diag.StringWriter) error {
s.Close()
return nil
},
@@ -87,9 +88,11 @@ func (s *session) handleRequests(in <-chan *ssh.Request, channel ssh.Channel) {
}
req.Reply(true, nil)
s.dispatchCommand(payload.Value, &stringWriter{channel})
dErr := s.commands.Dispatch(payload.Value, diag.NewWriter(channel))
status := struct{ Status uint32 }{uint32(0)}
// Report a real exit status rather than a hardcoded zero, so that
// `ssh nebula-host list-hostmap` is scriptable the same way `nebula ctl` is.
status := struct{ Status uint32 }{uint32(diag.StatusFor(dErr))}
channel.SendRequest("exit-status", false, ssh.Marshal(status))
channel.Close()
return
@@ -111,7 +114,7 @@ func (s *session) createTerm(channel ssh.Channel) *term.Terminal {
term.AutoCompleteCallback = func(line string, pos int, key rune) (newLine string, newPos int, ok bool) {
// key 9 is tab
if key == 9 {
cmds := matchCommand(s.commands, line)
cmds := s.commands.Match(line)
if len(cmds) == 1 {
return cmds[0] + " ", len(cmds[0]) + 1, true
}
@@ -128,49 +131,19 @@ func (s *session) createTerm(channel ssh.Channel) *term.Terminal {
}
func (s *session) handleInput() {
w := &stringWriter{w: s.term}
w := diag.NewWriter(s.term)
for {
line, err := s.term.ReadLine()
if err != nil {
break
}
s.dispatchCommand(line, w)
// The interactive console reports problems on the terminal the user is already
// looking at, so the error is nothing extra to say here.
_ = s.commands.Dispatch(line, w)
}
}
func (s *session) dispatchCommand(line string, w StringWriter) {
args, err := shlex.Split(line, true)
if err != nil {
return
}
if len(args) == 0 {
dumpCommands(s.commands, w)
return
}
c, err := lookupCommand(s.commands, args[0])
if err != nil {
return
}
if c == nil {
err := w.WriteLine(fmt.Sprintf("did not understand: %s", line))
_ = err
dumpCommands(s.commands, w)
return
}
if checkHelpArgs(args) {
s.dispatchCommand(fmt.Sprintf("%s %s", "help", c.Name), w)
return
}
_ = execCommand(c, args[1:], w)
}
func (s *session) Close() {
s.c.Close()
s.cancel()
+18 -5
View File
@@ -31,7 +31,8 @@ func procyield(cycles uint32)
const (
packetsPerRing = 1024
bytesPerPacket = 2048 - 32
// Caps tun.mtu at MTU-32 direct, MTU-64 relayed, unenforced anywhere else. 17.6MB page locked per socket.
bytesPerPacket = MTU
receiveSpins = 15
)
@@ -69,12 +70,14 @@ func NewRIOListener(l *slog.Logger, addr netip.Addr, port int) (*RIOConn, error)
err := u.bind(l, &windows.SockaddrInet6{Addr: addr.As16(), Port: port})
if err != nil {
u.close()
return nil, fmt.Errorf("bind: %w", err)
}
for i := 0; i < packetsPerRing; i++ {
err = u.insertReceiveRequest()
if err != nil {
u.close()
return nil, fmt.Errorf("init rx ring: %w", err)
}
}
@@ -356,15 +359,25 @@ func (u *RIOConn) Close() error {
return nil
}
u.close()
return nil
}
// Also unwinds a partial build from NewRIOListener, where isOpen is false and Close would no-op.
// Socket first, unlike wireguard-go: receive() re-arms every slot, so freeing the rings under a live socket
// hands the kernel freed pages for all packetsPerRing outstanding receives.
func (u *RIOConn) close() {
// WSASocket reports failure as InvalidHandle, not zero.
if u.sock != 0 && u.sock != windows.InvalidHandle {
windows.CloseHandle(u.sock)
}
u.sock = 0
windows.PostQueuedCompletionStatus(u.rx.iocp, 0, 0, nil)
windows.PostQueuedCompletionStatus(u.tx.iocp, 0, 0, nil)
u.rx.CloseAndZero()
u.tx.CloseAndZero()
if u.sock != 0 {
windows.CloseHandle(u.sock)
}
return nil
}
func (ring *ringBuffer) Push() *ringPacket {