Compare commits

..
Author SHA1 Message Date
Wade Simmons 7316283e0d for all of ./... 2026-09-01 13:01:40 -04:00
Wade Simmons ae0fa0f2af apply go fix for Go 1.26
This applies the `go fix` recommendations for go1.26
2026-09-01 11:39:02 -04:00
104 changed files with 1121 additions and 3463 deletions
+5 -34
View File
@@ -51,19 +51,15 @@ 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,$Ip6_1/64" -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" -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,$Ip6_2/64" -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" -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.
@@ -86,7 +82,7 @@ tun:
drop_local_broadcast: false
drop_multicast: false
tx_queue: 500
mtu: $Mtu
mtu: 1300
network_category: private
logging:
level: info
@@ -130,7 +126,7 @@ tun:
drop_local_broadcast: false
drop_multicast: false
tx_queue: 500
mtu: $Mtu
mtu: 1300
logging:
level: info
format: text
@@ -173,7 +169,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; { 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"
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"
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
@@ -218,16 +214,6 @@ 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"
@@ -235,13 +221,6 @@ 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"
@@ -255,14 +234,6 @@ 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.'
}
+1 -36
View File
@@ -7,32 +7,6 @@ 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
@@ -41,18 +15,11 @@ See the [v1.11.1](https://github.com/slackhq/nebula/milestone/30?closed=1) miles
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
@@ -917,9 +884,7 @@ created.)
- Initial public release.
[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
[Unreleased]: https://github.com/slackhq/nebula/compare/v1.10.3...HEAD
[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
+2 -11
View File
@@ -79,13 +79,7 @@ func (b *Bits) clearRange(startPos, count uint64) uint64 {
// handle the potential partial word before pos becomes u64 aligned
word := pos >> 6
bit := pos & 63
take := uint64(64) - bit
if take > remaining {
take = remaining
}
if take > b.length-pos {
take = b.length - pos
}
take := min(min(uint64(64)-bit, remaining), b.length-pos)
var mask uint64
if take == 64 {
mask = math.MaxUint64
@@ -189,10 +183,7 @@ func (b *Bits) Update(l *slog.Logger, i uint64) bool {
func (b *Bits) updateSlow(l *slog.Logger, i uint64) bool {
// If i is a jump, adjust the window, record lost, update current, and return true
if i > b.current {
end := i
if end > b.current+b.length {
end = b.current + b.length
}
end := min(i, b.current+b.length)
count := end - b.current
startPos := (b.current + 1) & b.lengthMask
+1 -1
View File
@@ -120,7 +120,7 @@ func TestCertificate_SignP256_AlwaysNormalized(t *testing.T) {
pub := elliptic.Marshal(elliptic.P256(), priv.PublicKey.X, priv.PublicKey.Y)
rawPriv := priv.D.FillBytes(make([]byte, 32))
for i := 0; i < 1000; i++ {
for i := range 1000 {
if i&1 == 1 {
tbs.Version = Version1
} else {
+4 -4
View File
@@ -151,7 +151,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
var groups []string
if *cf.groups != "" {
for _, rg := range strings.Split(*cf.groups, ",") {
for rg := range strings.SplitSeq(*cf.groups, ",") {
g := strings.TrimSpace(rg)
if g != "" {
groups = append(groups, g)
@@ -171,7 +171,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
}
if *cf.networks != "" {
for _, rs := range strings.Split(*cf.networks, ",") {
for rs := range strings.SplitSeq(*cf.networks, ",") {
rs := strings.Trim(rs, " ")
if rs != "" {
n, err := netip.ParsePrefix(rs)
@@ -193,7 +193,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
}
if *cf.unsafeNetworks != "" {
for _, rs := range strings.Split(*cf.unsafeNetworks, ",") {
for rs := range strings.SplitSeq(*cf.unsafeNetworks, ",") {
rs := strings.Trim(rs, " ")
if rs != "" {
n, err := netip.ParsePrefix(rs)
@@ -221,7 +221,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
if !isP11 && *cf.encryption {
passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE"))
if len(passphrase) == 0 {
for i := 0; i < 5; i++ {
for range 5 {
errOut.Write([]byte("Enter passphrase: "))
passphrase, err = pr.ReadPassword()
-1
View File
@@ -1,5 +1,4 @@
//go:build !windows
// +build !windows
package main
+4 -4
View File
@@ -146,7 +146,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE"))
if len(passphrase) == 0 {
// ask for a passphrase until we get one
for i := 0; i < 5; i++ {
for range 5 {
errOut.Write([]byte("Enter passphrase: "))
passphrase, err = pr.ReadPassword()
@@ -203,7 +203,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
}
if *sf.networks != "" {
for _, rs := range strings.Split(*sf.networks, ",") {
for rs := range strings.SplitSeq(*sf.networks, ",") {
rs := strings.Trim(rs, " ")
if rs != "" {
n, err := netip.ParsePrefix(rs)
@@ -228,7 +228,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
}
if *sf.unsafeNetworks != "" {
for _, rs := range strings.Split(*sf.unsafeNetworks, ",") {
for rs := range strings.SplitSeq(*sf.unsafeNetworks, ",") {
rs := strings.Trim(rs, " ")
if rs != "" {
n, err := netip.ParsePrefix(rs)
@@ -247,7 +247,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
var groups []string
if *sf.groups != "" {
for _, rg := range strings.Split(*sf.groups, ",") {
for rg := range strings.SplitSeq(*sf.groups, ",") {
g := strings.TrimSpace(rg)
if g != "" {
groups = append(groups, g)
-1
View File
@@ -1,5 +1,4 @@
//go:build !windows
// +build !windows
package main
-1
View File
@@ -1,5 +1,4 @@
//go:build !windows
// +build !windows
package main
-131
View File
@@ -1,131 +0,0 @@
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
@@ -1,70 +0,0 @@
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,26 +32,11 @@ 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 {
-1
View File
@@ -1,5 +1,4 @@
//go:build !linux
// +build !linux
package main
-922
View File
@@ -1,922 +0,0 @@
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
@@ -1,69 +0,0 @@
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)
}
})
}
+6 -9
View File
@@ -5,6 +5,7 @@ import (
"errors"
"fmt"
"log/slog"
"maps"
"math"
"os"
"os/signal"
@@ -154,9 +155,7 @@ func (c *C) ReloadConfig() {
defer c.reloadLock.Unlock()
c.oldSettings = make(map[string]any)
for k, v := range c.Settings {
c.oldSettings[k] = v
}
maps.Copy(c.oldSettings, c.Settings)
err := c.Load(c.path)
if err != nil {
@@ -177,9 +176,7 @@ func (c *C) ReloadConfigString(raw string) error {
defer c.reloadLock.Unlock()
c.oldSettings = make(map[string]any)
for k, v := range c.Settings {
c.oldSettings[k] = v
}
maps.Copy(c.oldSettings, c.Settings)
err := c.LoadString(raw)
if err != nil {
@@ -216,7 +213,7 @@ func (c *C) GetStringSlice(k string, d []string) []string {
}
v := make([]string, len(rv))
for i := 0; i < len(v); i++ {
for i := range v {
v[i] = fmt.Sprintf("%v", rv[i])
}
@@ -310,8 +307,8 @@ func (c *C) IsSet(k string) bool {
}
func (c *C) get(k string, v any) any {
parts := strings.Split(k, ".")
for _, p := range parts {
parts := strings.SplitSeq(k, ".")
for p := range parts {
m, ok := v.(map[string]any)
if !ok {
return nil
+1 -1
View File
@@ -97,7 +97,7 @@ func TestConnectionState_NextMessageCounter(t *testing.T) {
assert.Equal(t, RejectAfterMessages, cs.messageCounter.Load())
// Continued send attempts stay refused and the counter never wraps
for i := 0; i < 10; i++ {
for range 10 {
_, ok = cs.NextMessageCounter()
assert.False(t, ok)
}
-4
View File
@@ -50,7 +50,6 @@ type Control struct {
ctx context.Context
cancel context.CancelFunc
sshStart func()
ctlStart func()
statsStart func()
dnsStart func()
lighthouseStart func()
@@ -100,9 +99,6 @@ 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()
}
+1 -1
View File
@@ -247,7 +247,7 @@ func TestControl_ConcurrentStopAndStart(t *testing.T) {
c, _, _ := newReadyControl(t)
var wg sync.WaitGroup
for i := 0; i < 2; i++ {
for range 2 {
wg.Go(func() { c.Stop() })
}
wg.Go(func() { _ = c.Start() })
-234
View File
@@ -1,234 +0,0 @@
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
@@ -1,305 +0,0 @@
//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
@@ -1,41 +0,0 @@
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()
}
-226
View File
@@ -1,226 +0,0 @@
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
@@ -1,144 +0,0 @@
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
@@ -1,125 +0,0 @@
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
@@ -1,167 +0,0 @@
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
@@ -1,137 +0,0 @@
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
@@ -1,273 +0,0 @@
//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
@@ -1,107 +0,0 @@
//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
@@ -1,27 +0,0 @@
//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
}
-1
View File
@@ -1,5 +1,4 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
-1
View File
@@ -1,5 +1,4 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
-1
View File
@@ -1,5 +1,4 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
-7
View File
@@ -1,5 +1,4 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
@@ -116,9 +115,6 @@ 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",
@@ -216,9 +212,6 @@ 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",
-1
View File
@@ -1,5 +1,4 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
-1
View File
@@ -1,5 +1,4 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
-1
View File
@@ -1,5 +1,4 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
-1
View File
@@ -1,5 +1,4 @@
//go:build e2e_testing
// +build e2e_testing
package router
-1
View File
@@ -1,5 +1,4 @@
//go:build e2e_testing
// +build e2e_testing
package router
-1
View File
@@ -1,5 +1,4 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
-1
View File
@@ -1,5 +1,4 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
-24
View File
@@ -236,30 +236,6 @@ 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.
+1 -1
View File
@@ -22,7 +22,7 @@ func newFixedTicker(t *testing.T, l *slog.Logger, cacheLen int) *ConntrackCacheT
l: l,
cache: make(ConntrackCache, cacheLen),
}
for i := 0; i < cacheLen; i++ {
for i := range cacheLen {
c.cache[Packet{LocalPort: uint16(i) + 1}] = struct{}{}
}
c.cacheTick.Store(1) // cacheV starts at 0, so Get() takes the reset path
+1 -15
View File
@@ -34,9 +34,7 @@ type LightHouse struct {
myVpnNetworks []netip.Prefix
myVpnNetworksTable *bart.Lite
// myVpnAddrsTable contains our overlay host addrs, as opposed to the overlay networks
myVpnAddrsTable *bart.Lite
punchy *Punchy
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.
@@ -106,7 +104,6 @@ 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,
@@ -1161,17 +1158,6 @@ 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
+53 -79
View File
@@ -27,27 +27,15 @@ 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")
cs := testCertState(myVpnNet)
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
lh1 := "10.128.0.2"
c := config.NewC(l)
@@ -67,7 +55,12 @@ func Test_lhStaticMapping(t *testing.T) {
func TestReloadLighthouseInterval(t *testing.T) {
l := test.NewLogger()
myVpnNet := netip.MustParsePrefix("10.128.0.1/16")
cs := testCertState(myVpnNet)
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
lh1 := "10.128.0.2"
c := config.NewC(l)
@@ -97,7 +90,12 @@ func TestReloadLighthouseInterval(t *testing.T) {
func BenchmarkLighthouseHandleRequest(b *testing.B) {
l := test.NewLogger()
myVpnNet := netip.MustParsePrefix("10.128.0.1/0")
cs := testCertState(myVpnNet)
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
c := config.NewC(l)
lh, err := NewLightHouseFromConfig(b.Context(), l, c, cs, nil, nil)
@@ -197,7 +195,12 @@ func TestLighthouse_Memory(t *testing.T) {
c.Settings["listen"] = map[string]any{"port": 4242}
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
cs := testCertState(myVpnNet)
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
lh.ifce = &mockEncWriter{}
require.NoError(t, err)
@@ -277,7 +280,12 @@ func TestLighthouse_reload(t *testing.T) {
c.Settings["listen"] = map[string]any{"port": 4242}
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
cs := testCertState(myVpnNet)
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
require.NoError(t, err)
@@ -307,7 +315,12 @@ func TestLighthouse_reloadStaticHostMap(t *testing.T) {
}
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
cs := testCertState(myVpnNet)
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
require.NoError(t, err)
@@ -416,9 +429,7 @@ func TestLighthouse_reloadStaticHostMap(t *testing.T) {
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("3.3.3.3:4242")}, rl.CopyAddrs([]netip.Prefix{}))
}
// 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 {
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
req := &NebulaMeta{
Type: NebulaMeta_HostQuery,
Details: &NebulaMetaDetails{},
@@ -436,59 +447,12 @@ func sendLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr,
panic(err)
}
w := &testEncWriter{metaFilter: filter}
lhh.HandleRequest(fromAddr, []netip.Addr{myVpnIp}, b, w)
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"},
w := &testEncWriter{
metaFilter: &filter,
}
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")
lhh.HandleRequest(fromAddr, []netip.Addr{myVpnIp}, b, w)
return w.lastReply
}
func newLHHostUpdate(fromAddr netip.AddrPort, vpnIp netip.Addr, addrs []netip.AddrPort, lhh *LightHouseHandler) {
@@ -678,7 +642,12 @@ func TestLighthouse_Dont_Delete_Static_Hosts(t *testing.T) {
}
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
cs := testCertState(myVpnNet)
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
require.NoError(t, err)
lh.ifce = &mockEncWriter{}
@@ -739,7 +708,12 @@ func TestLighthouse_DeletesWork(t *testing.T) {
}
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
cs := testCertState(myVpnNet)
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
require.NoError(t, err)
lh.ifce = &mockEncWriter{}
+2 -13
View File
@@ -14,7 +14,6 @@ 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"
@@ -69,9 +68,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
}
l.Info("Firewall started", "firewallHashes", fw.GetRuleHashes())
commands := diag.NewRegistry()
ssh, err := sshd.NewSSHServer(ctx, l.With("subsystem", "sshd"), commands)
ssh, err := sshd.NewSSHServer(ctx, l.With("subsystem", "sshd"))
if err != nil {
return nil, util.ContextualizeIfNeeded("Error while creating SSH server", err)
}
@@ -325,20 +322,13 @@ 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, commands, ifce)
attachCommands(l, c, ssh, ifce)
networkChanges := udp.NewNetworkChangeMonitor(ctx, l, c)
@@ -349,7 +339,6 @@ 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,
-1
View File
@@ -1,5 +1,4 @@
//go:build boringcrypto
// +build boringcrypto
package noiseutil
-1
View File
@@ -1,5 +1,4 @@
//go:build boringcrypto
// +build boringcrypto
package noiseutil
+2 -2
View File
@@ -127,7 +127,7 @@ func TestPseudoSumIPv6MatchesReference(t *testing.T) {
func TestIPv4HdrChecksumMatchesReference(t *testing.T) {
rng := rand.New(rand.NewSource(0x1791))
for _, hdrLen := range []int{20, 24, 40, 60} {
for trial := 0; trial < 200; trial++ {
for trial := range 200 {
hdr := make([]byte, hdrLen)
rng.Read(hdr)
hdr[0] = 0x40 | byte(hdrLen/4)
@@ -156,7 +156,7 @@ func TestIPv4HdrChecksumMatchesReference(t *testing.T) {
// the way a receiver does (pseudo-header + L4 must sum to all-ones).
func TestChecksumSeedReceiverAcceptance(t *testing.T) {
rng := rand.New(rand.NewSource(0x1826))
for trial := 0; trial < 200; trial++ {
for trial := range 200 {
src := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
dst := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
payLen := rng.Intn(1500)
+1 -1
View File
@@ -540,7 +540,7 @@ func TestCoalescerCapBySegments(t *testing.T) {
c := newTestTCPCoalescer(t, w)
pay := make([]byte, 512)
seq := uint32(1000)
for i := 0; i < tcpCoalesceMaxSegs+5; i++ {
for range tcpCoalesceMaxSegs + 5 {
if err := c.Commit(buildTCPv4(seq, tcpAck, pay)); err != nil {
t.Fatal(err)
}
+2 -2
View File
@@ -28,7 +28,7 @@ func TestSendBatchReserveCommitFlush(t *testing.T) {
b := NewSendBatch(fw, 4, 32)
ap := netip.MustParseAddrPort("10.0.0.1:4242")
for i := 0; i < 4; i++ {
for i := range 4 {
slot := b.Reserve(32)
if cap(slot) != 32 {
t.Fatalf("slot %d: cap=%d want 32", i, cap(slot))
@@ -72,7 +72,7 @@ func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
b := NewSendBatch(fw, 3, 8)
ap := netip.MustParseAddrPort("10.0.0.1:80")
for i := 0; i < 3; i++ {
for i := range 3 {
s := b.Reserve(8)
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
b.Commit(pkt, ap)
+3 -3
View File
@@ -129,7 +129,7 @@ func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pay := make([]byte, 1200)
for i := 0; i < 3; i++ {
for range 3 {
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
t.Fatal(err)
}
@@ -259,7 +259,7 @@ func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pay := make([]byte, 100)
for i := 0; i < udpCoalesceMaxSegs+5; i++ {
for range udpCoalesceMaxSegs + 5 {
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
t.Fatal(err)
}
@@ -317,7 +317,7 @@ func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pay := make([]byte, 1200)
for i := 0; i < 3; i++ {
for range 3 {
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
t.Fatal(err)
}
+1 -1
View File
@@ -151,7 +151,7 @@ func TestChecksumTailPaths(t *testing.T) {
offsets := []int{0, 1, 3, 7, 15} // mix of aligned and odd starts
for k := 0; k <= maxK; k++ {
for tail := 0; tail < 64; tail++ {
for tail := range 64 {
length := 64*k + tail
for _, seed := range seeds {
for _, off := range offsets {
-1
View File
@@ -1,5 +1,4 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
-1
View File
@@ -1,5 +1,4 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
-1
View File
@@ -1,5 +1,4 @@
//go:build linux && !android
// +build linux,!android
package tio
-1
View File
@@ -1,5 +1,4 @@
//go:build linux && !android
// +build linux,!android
package tio
-1
View File
@@ -1,5 +1,4 @@
//go:build linux && !android
// +build linux,!android
package tio
-1
View File
@@ -1,5 +1,4 @@
//go:build linux && !android
// +build linux,!android
package tio
-1
View File
@@ -1,5 +1,4 @@
//go:build linux && !android
// +build linux,!android
package tio
+4 -7
View File
@@ -1,5 +1,4 @@
//go:build linux && !android && !e2e_testing
// +build linux,!android,!e2e_testing
package tio
@@ -117,17 +116,15 @@ func TestPoll_ConcurrentWrite_NoRace(t *testing.T) {
}()
var wg sync.WaitGroup
for w := 0; w < writers; w++ {
wg.Add(1)
go func() {
defer wg.Done()
for i := 0; i < perWriter; i++ {
for range writers {
wg.Go(func() {
for range perWriter {
if _, werr := p.Write(payload); werr != nil {
t.Errorf("write: %v", werr)
return
}
}
}()
})
}
wg.Wait()
-1
View File
@@ -1,5 +1,4 @@
//go:build linux && !android
// +build linux,!android
package tio
+5 -6
View File
@@ -1,5 +1,4 @@
//go:build linux && !android && !e2e_testing
// +build linux,!android,!e2e_testing
package tio
@@ -137,7 +136,7 @@ func buildTSOv4(t *testing.T, payLen, mss int) ([]byte, virtio.Hdr) {
binary.BigEndian.PutUint16(pkt[34:36], 65535) // window
// payload
for i := 0; i < payLen; i++ {
for i := range payLen {
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
}
return pkt, virtio.NewHeader(
@@ -257,7 +256,7 @@ func TestSegmentTCPv6(t *testing.T) {
pkt[53] = 0x19 // FIN | ACK | PSH — exercise FIN clearing too
binary.BigEndian.PutUint16(pkt[54:56], 65535)
for i := 0; i < payLen; i++ {
for i := range payLen {
pkt[ipLen+tcpLen+i] = byte(i)
}
@@ -361,7 +360,7 @@ func buildUSOv4(t *testing.T, payLen, gsoSize int) ([]byte, virtio.Hdr) {
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
binary.BigEndian.PutUint16(pkt[24:26], uint16(udpLen+payLen)) // superpacket length
for i := 0; i < payLen; i++ {
for i := range payLen {
pkt[ipLen+udpLen+i] = byte(i & 0xff)
}
@@ -473,7 +472,7 @@ func TestSegmentUDPv6(t *testing.T) {
// Superpacket-wide length, as the kernel supplies it; see buildUSOv4.
binary.BigEndian.PutUint16(pkt[44:46], uint16(udpLen+payLen))
for i := 0; i < payLen; i++ {
for i := range payLen {
pkt[ipLen+udpLen+i] = byte(i)
}
@@ -772,7 +771,7 @@ func buildTSOv6(payLen, gso int) []byte {
pkt[53] = 0x10 // ACK only
binary.BigEndian.PutUint16(pkt[54:56], 65535)
for i := 0; i < payLen; i++ {
for i := range payLen {
pkt[ipLen+tcpLen+i] = byte(i)
}
return pkt
-1
View File
@@ -1,5 +1,4 @@
//go:build linux && !android
// +build linux,!android
package virtio
+4 -11
View File
@@ -1,5 +1,4 @@
//go:build linux && !android
// +build linux,!android
// Package virtio implements the pure validation, header-correction, and
// per-segment slicing logic for kernel-supplied TSO/USO superpackets on
@@ -254,12 +253,9 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
var savedHdr [maxSegHdrLen]byte
copy(savedHdr[:headerLen], pkt[:headerLen])
for i := 0; i < numSeg; i++ {
for i := range numSeg {
segStart := i * gsoSize
segEnd := segStart + gsoSize
if segEnd > payLen {
segEnd = payLen
}
segEnd := min(segStart+gsoSize, payLen)
segPayLen := segEnd - segStart
segLen := headerLen + segPayLen
headerOff := i * gsoSize
@@ -359,12 +355,9 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
var savedHdr [maxSegHdrLen]byte
copy(savedHdr[:headerLen], pkt[:headerLen])
for i := 0; i < numSeg; i++ {
for i := range numSeg {
segStart := i * gsoSize
segEnd := segStart + gsoSize
if segEnd > payLen {
segEnd = payLen
}
segEnd := min(segStart+gsoSize, payLen)
segPayLen := segEnd - segStart
segLen := headerLen + segPayLen
headerOff := i * gsoSize
+7 -8
View File
@@ -1,5 +1,4 @@
//go:build linux && !android
// +build linux,!android
package virtio
@@ -58,7 +57,7 @@ func buildTCPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
pkt[33] = 0x18 // ACK | PSH
binary.BigEndian.PutUint16(pkt[34:36], 65535) // window
for i := 0; i < payLen; i++ {
for i := range payLen {
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
}
return pkt, ipLen + tcpLen, ipLen
@@ -82,7 +81,7 @@ func buildUDPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
for i := 0; i < payLen; i++ {
for i := range payLen {
pkt[ipLen+udpLen+i] = byte(i & 0xff)
}
return pkt, ipLen + udpLen, ipLen
@@ -188,7 +187,7 @@ func TestSegmentTCPHeaderNotCorrupted(t *testing.T) {
// Payload bytes must be the original contiguous slice.
segPayLen := len(seg) - int(hdrLen)
wantPay := make([]byte, segPayLen)
for k := 0; k < segPayLen; k++ {
for k := range segPayLen {
wantPay[k] = byte((off + k) & 0xff)
}
if !bytes.Equal(seg[hdrLen:], wantPay) {
@@ -317,7 +316,7 @@ func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
}
wantPay := make([]byte, segPayLen)
for k := 0; k < segPayLen; k++ {
for k := range segPayLen {
wantPay[k] = byte((off + k) & 0xff)
}
if !bytes.Equal(seg[hdrLen:], wantPay) {
@@ -365,7 +364,7 @@ func buildUDPv4Single(payload []byte) (pkt []byte, hdr Hdr) {
// 0xffff, because all-zero is the reserved "no checksum" encoding that IPv6 rejects outright.
func TestFinishChecksumUDPZeroStoresAllOnes(t *testing.T) {
var payload []byte
for i := 0; i < 0x10000; i++ {
for i := range 0x10000 {
p := []byte{byte(i >> 8), byte(i)}
pkt, hdr := buildUDPv4Single(p)
cs, co := int(hdr.CsumStart), int(hdr.CsumOffset)
@@ -544,7 +543,7 @@ func TestBaseSumsMatchZeroingReference(t *testing.T) {
t.Run("ipv4", func(t *testing.T) {
for ihl := ipv4HeaderMinLen; ihl <= ipv4HeaderMaxLen; ihl += 4 {
for iter := 0; iter < 5000; iter++ {
for range 5000 {
pkt := make([]byte, ihl)
for i := range pkt {
pkt[i] = randByte(&state)
@@ -574,7 +573,7 @@ func TestBaseSumsMatchZeroingReference(t *testing.T) {
for dataOff := 5; dataOff <= 15; dataOff++ {
tcpLen := dataOff * 4
headerLen := csumStart + tcpLen
for iter := 0; iter < 5000; iter++ {
for range 5000 {
pkt := make([]byte, headerLen+64)
for i := range pkt {
pkt[i] = randByte(&state)
+2 -2
View File
@@ -93,14 +93,14 @@ func prefixToMask(prefix netip.Prefix) netip.Addr {
}
func flipBytes(b []byte) []byte {
for i := 0; i < len(b); i++ {
for i := range b {
b[i] ^= 0xFF
}
return b
}
func orBytes(a []byte, b []byte) []byte {
ret := make([]byte, len(a))
for i := 0; i < len(a); i++ {
for i := range a {
ret[i] = a[i] | b[i]
}
return ret
-1
View File
@@ -1,5 +1,4 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
-1
View File
@@ -1,5 +1,4 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
-1
View File
@@ -1,5 +1,4 @@
//go:build !ios && !e2e_testing
// +build !ios,!e2e_testing
package overlay
-1
View File
@@ -1,5 +1,4 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
-1
View File
@@ -1,5 +1,4 @@
//go:build ios && !e2e_testing
// +build ios,!e2e_testing
package overlay
+1 -2
View File
@@ -1,5 +1,4 @@
//go:build !android && !e2e_testing
// +build !android,!e2e_testing
package overlay
@@ -548,7 +547,7 @@ func (t *tun) setDefaultRoute(cidr netip.Prefix) error {
if err != nil {
t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr)
//retry twice more -- on some systems there appears to be a race condition where if we set routes too soon, netlink says `invalid argument`
for i := 0; i < 2; i++ {
for range 2 {
time.Sleep(100 * time.Millisecond)
err = netlink.RouteReplace(&nr)
if err == nil {
-1
View File
@@ -1,5 +1,4 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
-1
View File
@@ -1,5 +1,4 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
-1
View File
@@ -1,5 +1,4 @@
//go:build !windows
// +build !windows
package overlay
-1
View File
@@ -1,5 +1,4 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
-1
View File
@@ -1,5 +1,4 @@
//go:build e2e_testing
// +build e2e_testing
package overlay
-30
View File
@@ -1,5 +1,4 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
@@ -11,7 +10,6 @@ import (
"os"
"path/filepath"
"runtime"
"slices"
"sync/atomic"
"syscall"
"unsafe"
@@ -183,7 +181,6 @@ 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 {
@@ -191,9 +188,6 @@ 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 {
@@ -215,11 +209,6 @@ 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)
@@ -234,25 +223,6 @@ func (t *winTun) setMTU(luid winipcfg.LUID, foundDefault4, carriesV6 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
}
+2 -2
View File
@@ -108,7 +108,7 @@ func TestUserDeviceReadersConcurrentRace(t *testing.T) {
var wg sync.WaitGroup
run := func(idx int) {
defer wg.Done()
for i := 0; i < iterations; i++ {
for range iterations {
pkts, err := readers[idx].Read()
if err != nil {
errs <- err
@@ -136,7 +136,7 @@ func TestUserDeviceReadersConcurrentRace(t *testing.T) {
// waiting reader's private buffer, so reusing buf between writes is safe.
go func() {
buf := make([]byte, 32)
for i := 0; i < 2*iterations; i++ {
for i := range 2 * iterations {
for j := range buf {
buf[j] = byte(i + j)
}
+3 -3
View File
@@ -27,7 +27,7 @@ func TestPacketsAreBalancedEqually(t *testing.T) {
gw3count := 0
iterationCount := uint16(65535)
for i := uint16(0); i < iterationCount; i++ {
for i := range iterationCount {
packet := firewall.Packet{
LocalAddr: netip.MustParseAddr("192.168.1.1"),
RemoteAddr: netip.MustParseAddr("10.0.0.1"),
@@ -74,7 +74,7 @@ func TestPacketsAreBalancedByPriority(t *testing.T) {
gw2count := 0
iterationCount := uint16(65535)
for i := uint16(0); i < iterationCount; i++ {
for i := range iterationCount {
packet := firewall.Packet{
LocalAddr: netip.MustParseAddr("192.168.1.1"),
RemoteAddr: netip.MustParseAddr("10.0.0.1"),
@@ -115,7 +115,7 @@ func TestBalancePacketDistributsRandomlyAndReturnsFalseIfBucketsNotCalculated(t
gw1count := 0
gw2count := 0
for i := uint16(0); i < iterationCount; i++ {
for i := range iterationCount {
packet := firewall.Packet{
LocalAddr: netip.MustParseAddr("192.168.1.1"),
RemoteAddr: netip.MustParseAddr("10.0.0.1"),
+5 -4
View File
@@ -3,6 +3,7 @@ package routing
import (
"fmt"
"net/netip"
"strings"
)
const (
@@ -13,14 +14,14 @@ const (
type Gateways []Gateway
func (g Gateways) String() string {
str := ""
var str strings.Builder
for i, gw := range g {
str += gw.String()
str.WriteString(gw.String())
if i < len(g)-1 {
str += ", "
str.WriteString(", ")
}
}
return str
return str.String()
}
type Gateway struct {
+4 -8
View File
@@ -1,21 +1,19 @@
package nebula
import (
"context"
"testing"
"time"
)
func TestScheduler_PooledReuse(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
ctx := t.Context()
s := NewScheduler[int](16)
delivered := make(chan int, 256)
go s.Run(ctx, func(item int) { delivered <- item })
const N = 100
for i := 0; i < N; i++ {
for i := range N {
s.Schedule(ctx, i, time.Millisecond)
}
@@ -34,8 +32,7 @@ func TestScheduler_PooledReuse(t *testing.T) {
// BenchmarkScheduler_Schedule reports allocations per Schedule call.
// In steady state the Scheduler's sync.Pool means we should see zero allocs per op once the pool warms up.
func BenchmarkScheduler_Schedule(b *testing.B) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
ctx := b.Context()
s := NewScheduler[int](b.N)
go s.Run(ctx, func(int) {})
@@ -51,8 +48,7 @@ func BenchmarkScheduler_Schedule(b *testing.B) {
// What we'd pay per Schedule if Punchy called time.AfterFunc directly without the pooled Scheduler.
// Allocates a *time.Timer plus a closure each call.
func BenchmarkBareAfterFunc(b *testing.B) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
ctx := b.Context()
queue := make(chan int, b.N)
go func() {
+905 -3
View File
@@ -1,19 +1,62 @@
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) {
@@ -154,3 +197,862 @@ 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
}
+3 -16
View File
@@ -1,4 +1,4 @@
package diag
package sshd
import (
"errors"
@@ -10,17 +10,6 @@ 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)
@@ -55,10 +44,8 @@ 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, 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)
// fl.Parse has dumped error information to the user via the w writer.
return err
}
args = fl.Args()
}
+20 -7
View File
@@ -9,9 +9,8 @@ import (
"net"
"sync"
"github.com/armon/go-radix"
"golang.org/x/crypto/ssh"
"github.com/slackhq/nebula/diag"
)
type SSHServer struct {
@@ -26,9 +25,10 @@ type SSHServer struct {
trustedKeys map[string]map[string]bool
trustedCAs []ssh.PublicKey
// The commands this server serves. Shared with every other transport, see diag.Registry.
commands *diag.Registry
listener net.Listener
// List of available commands
helpCommand *Command
commands *radix.Tree
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, commands *diag.Registry) (*SSHServer, error) {
func NewSSHServer(ctx context.Context, l *slog.Logger) (*SSHServer, error) {
s := &SSHServer{
trustedKeys: make(map[string]map[string]bool),
l: l,
commands: commands,
commands: radix.New(),
ctx: ctx,
}
@@ -90,6 +90,14 @@ func NewSSHServer(ctx context.Context, l *slog.Logger, commands *diag.Registry)
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
}
@@ -152,6 +160,11 @@ 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.
+45 -18
View File
@@ -1,38 +1,37 @@
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 *diag.Registry
commands *radix.Tree
cancel func()
}
func NewSession(commands *diag.Registry, conn *ssh.ServerConn, chans <-chan ssh.NewChannel, cancel func(), l *slog.Logger) *session {
func NewSession(commands *radix.Tree, conn *ssh.ServerConn, chans <-chan ssh.NewChannel, cancel func(), l *slog.Logger) *session {
s := &session{
// A copy, so the logout command this session adds for itself stays invisible to every
// other session and to `nebula ctl`.
commands: commands.Clone(),
commands: radix.NewFromMap(commands.ToMap()),
l: l,
c: conn,
cancel: cancel,
}
s.commands.RegisterCommand(&diag.Command{
s.commands.Insert("logout", &Command{
Name: "logout",
ShortDescription: "Ends the current session",
Callback: func(a any, args []string, w diag.StringWriter) error {
Callback: func(a any, args []string, w StringWriter) error {
s.Close()
return nil
},
@@ -88,11 +87,9 @@ func (s *session) handleRequests(in <-chan *ssh.Request, channel ssh.Channel) {
}
req.Reply(true, nil)
dErr := s.commands.Dispatch(payload.Value, diag.NewWriter(channel))
s.dispatchCommand(payload.Value, &stringWriter{channel})
// 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))}
status := struct{ Status uint32 }{uint32(0)}
channel.SendRequest("exit-status", false, ssh.Marshal(status))
channel.Close()
return
@@ -114,7 +111,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 := s.commands.Match(line)
cmds := matchCommand(s.commands, line)
if len(cmds) == 1 {
return cmds[0] + " ", len(cmds[0]) + 1, true
}
@@ -131,19 +128,49 @@ func (s *session) createTerm(channel ssh.Channel) *term.Terminal {
}
func (s *session) handleInput() {
w := diag.NewWriter(s.term)
w := &stringWriter{w: s.term}
for {
line, err := s.term.ReadLine()
if err != nil {
break
}
// 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)
s.dispatchCommand(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()
+1 -7
View File
@@ -1,4 +1,4 @@
package diag
package sshd
import "io"
@@ -30,9 +30,3 @@ 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}
}
+2 -2
View File
@@ -25,7 +25,7 @@ func AssertDeepCopyEqual(t *testing.T, a any, b any) {
}
func traverseDeepCopy(t *testing.T, v1 reflect.Value, v2 reflect.Value, name string) bool {
if v1.Type() == v2.Type() && v1.Type() == reflect.TypeOf(netip.Addr{}) {
if v1.Type() == v2.Type() && v1.Type() == reflect.TypeFor[netip.Addr]() {
// Ignore netip.Addr types since they reuse an interned global value
return false
}
@@ -72,7 +72,7 @@ func traverseDeepCopy(t *testing.T, v1 reflect.Value, v2 reflect.Value, name str
}
return traverseDeepCopy(t, v1.Elem(), v2.Elem(), name)
case reflect.Ptr:
case reflect.Pointer:
local := reflect.ValueOf(time.Local).Pointer()
if local == v1.Pointer() && local == v2.Pointer() {
return true
-1
View File
@@ -1,5 +1,4 @@
//go:build darwin && !ios && !e2e_testing
// +build darwin,!ios,!e2e_testing
package udp
-1
View File
@@ -1,5 +1,4 @@
//go:build darwin && !ios && !e2e_testing
// +build darwin,!ios,!e2e_testing
package udp
-1
View File
@@ -1,5 +1,4 @@
//go:build !darwin || ios || e2e_testing
// +build !darwin ios e2e_testing
package udp
-1
View File
@@ -1,5 +1,4 @@
//go:build !e2e_testing
// +build !e2e_testing
package udp
-1
View File
@@ -1,5 +1,4 @@
//go:build !e2e_testing
// +build !e2e_testing
package udp
-1
View File
@@ -1,5 +1,4 @@
//go:build !e2e_testing
// +build !e2e_testing
package udp
+2 -5
View File
@@ -294,10 +294,7 @@ func deliverSegments(r EncReader, from netip.AddrPort, payload []byte, segSize i
return
}
for off := 0; off < len(payload); off += segSize {
end := off + segSize
if end > len(payload) {
end = len(payload)
}
end := min(off+segSize, len(payload))
r(from, payload[off:end:end])
}
}
@@ -479,7 +476,7 @@ func NewUDPStatsEmitter(udpConns []Conn) func() {
return func() {
for i, gauges := range udpGauges {
if err := udpConns[i].(*StdConn).getMemInfo(&meminfo); err == nil {
for j := 0; j < unix.SK_MEMINFO_VARS; j++ {
for j := range unix.SK_MEMINFO_VARS {
gauges[j].Update(int64(meminfo[j]))
}
}
+4 -4
View File
@@ -107,7 +107,7 @@ func TestWriteBatchBadFamilyDeliversOthers(t *testing.T) {
got := map[string]bool{}
rx.SetReadDeadline(time.Now().Add(2 * time.Second))
buf := make([]byte, 64)
for i := 0; i < 2; i++ {
for i := range 2 {
n, _, rerr := rx.ReadFromUDPAddrPort(buf)
if rerr != nil {
t.Fatalf("expected 2 delivered packets, read #%d failed: %v", i+1, rerr)
@@ -161,7 +161,7 @@ func TestWriteBatchUnreachableDestDeliversOthers(t *testing.T) {
got := map[string]bool{}
rx.SetReadDeadline(time.Now().Add(2 * time.Second))
buf := make([]byte, 64)
for i := 0; i < 4; i++ {
for i := range 4 {
n, _, rerr := rx.ReadFromUDPAddrPort(buf)
if rerr != nil {
t.Fatalf("expected 4 delivered packets, read #%d failed: %v (got so far: %v)", i+1, rerr, got)
@@ -691,7 +691,7 @@ func TestGSOEngagesOnLoopback(t *testing.T) {
// The kernel must deliver the original datagram boundaries and bytes.
_ = rx.SetReadDeadline(time.Now().Add(5 * time.Second))
got := make([]byte, pktLen+1)
for i := 0; i < numPkts; i++ {
for i := range numPkts {
n, _, err := rx.ReadFromUDP(got)
if err != nil {
t.Fatalf("rx read %d: %v", i, err)
@@ -699,7 +699,7 @@ func TestGSOEngagesOnLoopback(t *testing.T) {
if n != pktLen {
t.Fatalf("rx read %d: len=%d want %d (kernel segmented at wrong boundary)", i, n, pktLen)
}
for j := 0; j < n; j++ {
for j := range n {
if got[j] != byte(i) {
t.Fatalf("rx read %d: byte %d = %#x, want %#x", i, j, got[j], byte(i))
}
+3 -6
View File
@@ -112,7 +112,7 @@ func (w *batchWriter) prepareWriteMessages(n int, offloadsEnabled bool) {
w.cmsg = make([]byte, n*w.cmsgSpace)
for k := 0; k < n; k++ {
for k := range n {
base := k * w.cmsgSpace
seg := (*unix.Cmsghdr)(unsafe.Pointer(&w.cmsg[base]))
seg.Level = unix.SOL_UDP
@@ -214,7 +214,7 @@ func (w *batchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, er
break
}
for k := 0; k < runLen; k++ {
for k := range runLen {
b := bufs[i+k]
if len(b) == 0 {
w.iovs[iovIdx+k].Base = nil
@@ -318,10 +318,7 @@ func (w *batchWriter) planRun(bufs [][]byte, addrs []netip.AddrPort, start, iovB
return 1, segSize
}
dst := addrs[start]
maxLen := w.maxGSOSegments
if iovBudget < maxLen {
maxLen = iovBudget
}
maxLen := min(iovBudget, w.maxGSOSegments)
runLen := 1
total := segSize
for runLen < maxLen && start+runLen < len(bufs) {
+2 -2
View File
@@ -64,13 +64,13 @@ func TestWriteBatchNoAllocs(t *testing.T) {
addrs = append(addrs, dst)
}
// GSO-eligible run with a short tail.
for k := 0; k < 8; k++ {
for range 8 {
add(payload, dstA)
}
add(short, dstA)
add(payload, dstA)
// Alternating destinations defeat coalescing entirely.
for k := 0; k < 4; k++ {
for k := range 4 {
dst := dstA
if k%2 == 0 {
dst = dstB
-1
View File
@@ -1,5 +1,4 @@
//go:build !e2e_testing
// +build !e2e_testing
package udp

Some files were not shown because too many files have changed in this diff Show More