mirror of
https://github.com/slackhq/nebula.git
synced 2026-09-30 06:46:38 +02:00
Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
71c5297134 |
@@ -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
@@ -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
|
||||
|
||||
@@ -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]
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
-922
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -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,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
@@ -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) }
|
||||
@@ -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
@@ -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])
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
})
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
})
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -116,9 +116,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 +213,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",
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -807,7 +807,7 @@ func (hm *HandshakeManager) beginHandshake(via ViaSender, packet []byte, h *head
|
||||
}
|
||||
|
||||
hm.sendHandshakeResponse(via, response, hostinfo, false)
|
||||
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
||||
hostinfo.remotes.RefreshFromHandshake(vpnAddrs, remoteCert.Certificate.Version())
|
||||
|
||||
// Don't wait for UpdateWorker
|
||||
if f.lightHouse.IsAnyLighthouseAddr(vpnAddrs) {
|
||||
@@ -995,7 +995,7 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
||||
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
||||
}
|
||||
|
||||
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
||||
hostinfo.remotes.RefreshFromHandshake(vpnAddrs, remoteCert.Certificate.Version())
|
||||
f.metricHandshakes.Update(duration)
|
||||
|
||||
// Don't wait for UpdateWorker
|
||||
|
||||
+51
-32
@@ -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,
|
||||
@@ -516,17 +513,15 @@ func (lh *LightHouse) QueryServer(vpnAddr netip.Addr) {
|
||||
}
|
||||
|
||||
func (lh *LightHouse) QueryCache(vpnAddrs []netip.Addr) *RemoteList {
|
||||
lh.RLock()
|
||||
if v, ok := lh.addrMap[vpnAddrs[0]]; ok {
|
||||
lh.RUnlock()
|
||||
return v
|
||||
rl, ok := lh.findRemoteList(vpnAddrs)
|
||||
if ok {
|
||||
return rl
|
||||
}
|
||||
lh.RUnlock()
|
||||
|
||||
lh.Lock()
|
||||
defer lh.Unlock()
|
||||
// Add an entry if we don't already have one
|
||||
return lh.unlockedGetRemoteList(vpnAddrs) //todo CERT-V2 this contains addrmap lookups we could potentially skip
|
||||
return lh.unlockedGetRemoteList(vpnAddrs) //todo this re-calls unlockedFindRemoteList
|
||||
}
|
||||
|
||||
// queryAndPrepMessage is a lock helper on RemoteList, assisting the caller to build a lighthouse message containing
|
||||
@@ -672,23 +667,58 @@ func (lh *LightHouse) addCalculatedRemotes(vpnAddr netip.Addr) bool {
|
||||
return len(calculatedV4) > 0 || len(calculatedV6) > 0
|
||||
}
|
||||
|
||||
func (lh *LightHouse) findRemoteList(vpnAddrs []netip.Addr) (*RemoteList, bool) {
|
||||
lh.RLock()
|
||||
defer lh.RUnlock()
|
||||
return lh.unlockedFindRemoteList(vpnAddrs)
|
||||
}
|
||||
|
||||
// unlockedFindRemoteList checks addrMap for each of vpnAddrs. It returns the first RemoteList found,
|
||||
// and true if that RemoteList is present for all vpnAddrs.
|
||||
// If false, it means the addrMap, and possibly the RemoteList, need to be corrected.
|
||||
func (lh *LightHouse) unlockedFindRemoteList(vpnAddrs []netip.Addr) (*RemoteList, bool) {
|
||||
var am *RemoteList
|
||||
//todo: if a host with addresses A and B is "split", so it only has address A and a new host has only address B
|
||||
//todo: we don't handly that correctly, I'm pretty sure.
|
||||
missingOrDifferent := false
|
||||
for _, addr := range vpnAddrs {
|
||||
found, ok := lh.addrMap[addr]
|
||||
if !ok {
|
||||
missingOrDifferent = true
|
||||
} else if am == nil {
|
||||
am = found //the first list we find wins
|
||||
} else if am != found {
|
||||
missingOrDifferent = true
|
||||
}
|
||||
}
|
||||
return am, !missingOrDifferent
|
||||
}
|
||||
|
||||
// unlockedGetRemoteList assumes you have the lh lock
|
||||
func (lh *LightHouse) unlockedGetRemoteList(allAddrs []netip.Addr) *RemoteList {
|
||||
// before we go and make a new remotelist, we need to make sure we don't have one for any of this set of vpnaddrs yet
|
||||
for i, addr := range allAddrs {
|
||||
am, ok := lh.addrMap[addr]
|
||||
if ok {
|
||||
if i != 0 {
|
||||
lh.addrMap[allAddrs[0]] = am
|
||||
}
|
||||
return am
|
||||
am, ok := lh.unlockedFindRemoteList(allAddrs)
|
||||
|
||||
// we failed to find any RemoteLists: make a new one, fill out the addrMap
|
||||
if am == nil {
|
||||
am = NewRemoteList(allAddrs, lh.shouldAdd)
|
||||
for _, addr := range allAddrs {
|
||||
lh.addrMap[addr] = am
|
||||
}
|
||||
return am
|
||||
}
|
||||
|
||||
// we found one! Do we need to fix it?
|
||||
if !ok {
|
||||
am.Lock()
|
||||
am.vpnAddrs = make([]netip.Addr, len(allAddrs))
|
||||
copy(am.vpnAddrs, allAddrs)
|
||||
am.Unlock()
|
||||
for _, addr := range allAddrs {
|
||||
lh.addrMap[addr] = am
|
||||
}
|
||||
}
|
||||
|
||||
am := NewRemoteList(allAddrs, lh.shouldAdd)
|
||||
for _, addr := range allAddrs {
|
||||
lh.addrMap[addr] = am
|
||||
}
|
||||
return am
|
||||
}
|
||||
|
||||
@@ -1161,17 +1191,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
|
||||
|
||||
+173
-79
@@ -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{}
|
||||
@@ -764,3 +738,123 @@ func TestLighthouse_DeletesWork(t *testing.T) {
|
||||
out = lh.Query(testHost)
|
||||
assert.Nil(t, out)
|
||||
}
|
||||
|
||||
// newLHHostUpdateV2 sends a v2-style HostUpdateNotification where the sending tunnel carries
|
||||
// multiple vpn addrs (a dual-stack v2 cert). Details.VpnAddr is left blank like SendUpdate does.
|
||||
func newLHHostUpdateV2(fromAddr netip.AddrPort, vpnAddrs []netip.Addr, addrs []netip.AddrPort, lhh *LightHouseHandler) {
|
||||
req := &NebulaMeta{
|
||||
Type: NebulaMeta_HostUpdateNotification,
|
||||
Details: &NebulaMetaDetails{},
|
||||
}
|
||||
for _, v := range addrs {
|
||||
if v.Addr().Is4() {
|
||||
req.Details.V4AddrPorts = append(req.Details.V4AddrPorts, netAddrToProtoV4AddrPort(v.Addr(), v.Port()))
|
||||
} else {
|
||||
req.Details.V6AddrPorts = append(req.Details.V6AddrPorts, netAddrToProtoV6AddrPort(v.Addr(), v.Port()))
|
||||
}
|
||||
}
|
||||
b, err := req.Marshal()
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
lhh.HandleRequest(fromAddr, vpnAddrs, b, &testEncWriter{})
|
||||
}
|
||||
|
||||
func newIssue1868Lighthouse(t *testing.T) (*LightHouse, *LightHouseHandler) {
|
||||
l := test.NewLogger()
|
||||
c := config.NewC(l)
|
||||
c.Settings["lighthouse"] = map[string]any{"am_lighthouse": true}
|
||||
c.Settings["listen"] = map[string]any{"port": 4242}
|
||||
|
||||
myVpnNet4 := netip.MustParsePrefix("10.128.0.1/24")
|
||||
myVpnNet6 := netip.MustParsePrefix("fd00::1/64")
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet4)
|
||||
nt.Insert(myVpnNet6)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet4, myVpnNet6},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||
require.NoError(t, err)
|
||||
lh.ifce = &mockEncWriter{}
|
||||
return lh, lh.NewRequestHandler()
|
||||
}
|
||||
|
||||
// Scenario A: host registers via v2 (both addrs), then a v1 handshake with the same host completes on
|
||||
// the lighthouse (rehandshake after cert renewal, relay-initiated handshake, traffic to the LH's v4
|
||||
// addr...). handshake_manager does QueryCache(vpnAddrs) + RefreshFromHandshake(vpnAddrs) with the v1
|
||||
// cert's single address, which truncates RemoteList.vpnAddrs.
|
||||
func TestLighthouse_Issue1868_V1HandshakeTruncatesVpnAddrs(t *testing.T) {
|
||||
lh, lhh := newIssue1868Lighthouse(t)
|
||||
|
||||
hostV4 := netip.MustParseAddr("10.128.0.3")
|
||||
hostV6 := netip.MustParseAddr("fd00::3")
|
||||
hostUdp := netip.MustParseAddrPort("192.0.2.3:4242")
|
||||
hostLan := netip.MustParseAddrPort("10.0.0.3:4242")
|
||||
|
||||
askerV4 := netip.MustParseAddr("10.128.0.2")
|
||||
askerUdp := netip.MustParseAddrPort("192.0.2.2:4242")
|
||||
|
||||
// Boot: host handshakes with the LH using its v2 cert and sends an update
|
||||
newLHHostUpdateV2(hostUdp, []netip.Addr{hostV4, hostV6}, []netip.AddrPort{hostLan}, lhh)
|
||||
|
||||
// Both addresses resolve
|
||||
r := newLHHostRequest(askerUdp, askerV4, hostV4, lhh)
|
||||
require.NotNil(t, r.msg, "v4 query should be answered")
|
||||
assertIp4InArray(t, r.msg.Details.V4AddrPorts, hostLan)
|
||||
|
||||
r = newLHHostRequest(askerUdp, askerV4, hostV6, lhh)
|
||||
require.NotNil(t, r.msg, "v6 query should be answered before the v1 handshake")
|
||||
assertIp4InArray(t, r.msg.Details.V4AddrPorts, hostLan)
|
||||
|
||||
// Later: a v1 handshake with the same host completes on the LH. This is exactly what
|
||||
// handshake_manager.go does on completion, with the v1 cert's single vpn addr.
|
||||
rl := lh.QueryCache([]netip.Addr{hostV4})
|
||||
rl.RefreshFromHandshake([]netip.Addr{hostV4}, cert.Version1)
|
||||
|
||||
// The host keeps sending v2 updates over its v2 tunnel too
|
||||
newLHHostUpdateV2(hostUdp, []netip.Addr{hostV4, hostV6}, []netip.AddrPort{hostLan}, lhh)
|
||||
|
||||
r = newLHHostRequest(askerUdp, askerV4, hostV4, lhh)
|
||||
require.NotNil(t, r.msg, "v4 query should still be answered")
|
||||
assertIp4InArray(t, r.msg.Details.V4AddrPorts, hostLan)
|
||||
|
||||
r = newLHHostRequest(askerUdp, askerV4, hostV6, lhh)
|
||||
if assert.NotNil(t, r.msg, "BUG: v6 query is silently dropped after a v1 handshake truncated RemoteList.vpnAddrs") {
|
||||
assertIp4InArray(t, r.msg.Details.V4AddrPorts, hostLan)
|
||||
}
|
||||
}
|
||||
|
||||
// Scenario B: the LH first creates a RemoteList for the host keyed only by its v4 addr (a pending
|
||||
// LH-initiated v1 handshake does QueryCache([v4]) in handleOutbound), then the host arrives with v2.
|
||||
// unlockedGetRemoteList/QueryCache hit on allAddrs[0] and never add the v6 key to addrMap.
|
||||
func TestLighthouse_Issue1868_V4OnlyListNeverGainsV6Key(t *testing.T) {
|
||||
lh, lhh := newIssue1868Lighthouse(t)
|
||||
|
||||
hostV4 := netip.MustParseAddr("10.128.0.3")
|
||||
hostV6 := netip.MustParseAddr("fd00::3")
|
||||
hostUdp := netip.MustParseAddrPort("192.0.2.3:4242")
|
||||
hostLan := netip.MustParseAddrPort("10.0.0.3:4242")
|
||||
|
||||
askerV4 := netip.MustParseAddr("10.128.0.2")
|
||||
askerUdp := netip.MustParseAddrPort("192.0.2.2:4242")
|
||||
|
||||
// LH is a relay and someone asked it to relay to hostV4 while the host was offline:
|
||||
// StartHandshake(hostV4) -> handleOutbound -> QueryCache([hostV4]) creates a v4-only list.
|
||||
_ = lh.QueryCache([]netip.Addr{hostV4})
|
||||
|
||||
// Host boots and handshakes v2 with the LH (responder path does QueryCache + RefreshFromHandshake)
|
||||
rl := lh.QueryCache([]netip.Addr{hostV4, hostV6})
|
||||
rl.RefreshFromHandshake([]netip.Addr{hostV4, hostV6}, cert.Version2)
|
||||
newLHHostUpdateV2(hostUdp, []netip.Addr{hostV4, hostV6}, []netip.AddrPort{hostLan}, lhh)
|
||||
|
||||
r := newLHHostRequest(askerUdp, askerV4, hostV4, lhh)
|
||||
require.NotNil(t, r.msg, "v4 query should be answered")
|
||||
assertIp4InArray(t, r.msg.Details.V4AddrPorts, hostLan)
|
||||
|
||||
r = newLHHostRequest(askerUdp, askerV4, hostV6, lhh)
|
||||
if assert.NotNil(t, r.msg, "BUG: v6 query is silently dropped, addrMap never got the v6 key") {
|
||||
assertIp4InArray(t, r.msg.Details.V4AddrPorts, hostLan)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -11,7 +11,6 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"slices"
|
||||
"sync/atomic"
|
||||
"syscall"
|
||||
"unsafe"
|
||||
@@ -183,7 +182,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 +189,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 +210,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 +224,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
|
||||
}
|
||||
|
||||
|
||||
+9
-3
@@ -11,6 +11,8 @@ import (
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula/cert"
|
||||
)
|
||||
|
||||
// forEachFunc is used to benefit folks that want to do work inside the lock
|
||||
@@ -408,11 +410,15 @@ func (r *RemoteList) CopyBlockedRemotes() []netip.AddrPort {
|
||||
}
|
||||
|
||||
// RefreshFromHandshake locks and updates the RemoteList to account for data learned upon a completed handshake
|
||||
func (r *RemoteList) RefreshFromHandshake(vpnAddrs []netip.Addr) {
|
||||
func (r *RemoteList) RefreshFromHandshake(vpnAddrs []netip.Addr, v cert.Version) {
|
||||
r.Lock()
|
||||
r.badRemotes = nil
|
||||
r.vpnAddrs = make([]netip.Addr, len(vpnAddrs))
|
||||
copy(r.vpnAddrs, vpnAddrs)
|
||||
if v != cert.Version1 {
|
||||
// a handshake from a v1 cert can never expand our knowledge of the number of addresses a host has,
|
||||
// and, because v2 certs exist, it can also never contract it. So, only update this for non-v1 certs.
|
||||
r.vpnAddrs = make([]netip.Addr, len(vpnAddrs))
|
||||
copy(r.vpnAddrs, vpnAddrs)
|
||||
}
|
||||
r.Unlock()
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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,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}
|
||||
}
|
||||
+5
-18
@@ -31,8 +31,7 @@ func procyield(cycles uint32)
|
||||
|
||||
const (
|
||||
packetsPerRing = 1024
|
||||
// Caps tun.mtu at MTU-32 direct, MTU-64 relayed, unenforced anywhere else. 17.6MB page locked per socket.
|
||||
bytesPerPacket = MTU
|
||||
bytesPerPacket = 2048 - 32
|
||||
receiveSpins = 15
|
||||
)
|
||||
|
||||
@@ -70,14 +69,12 @@ func NewRIOListener(l *slog.Logger, addr netip.Addr, port int) (*RIOConn, error)
|
||||
|
||||
err := u.bind(l, &windows.SockaddrInet6{Addr: addr.As16(), Port: port})
|
||||
if err != nil {
|
||||
u.close()
|
||||
return nil, fmt.Errorf("bind: %w", err)
|
||||
}
|
||||
|
||||
for i := 0; i < packetsPerRing; i++ {
|
||||
err = u.insertReceiveRequest()
|
||||
if err != nil {
|
||||
u.close()
|
||||
return nil, fmt.Errorf("init rx ring: %w", err)
|
||||
}
|
||||
}
|
||||
@@ -359,25 +356,15 @@ func (u *RIOConn) Close() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
u.close()
|
||||
return nil
|
||||
}
|
||||
|
||||
// Also unwinds a partial build from NewRIOListener, where isOpen is false and Close would no-op.
|
||||
// Socket first, unlike wireguard-go: receive() re-arms every slot, so freeing the rings under a live socket
|
||||
// hands the kernel freed pages for all packetsPerRing outstanding receives.
|
||||
func (u *RIOConn) close() {
|
||||
// WSASocket reports failure as InvalidHandle, not zero.
|
||||
if u.sock != 0 && u.sock != windows.InvalidHandle {
|
||||
windows.CloseHandle(u.sock)
|
||||
}
|
||||
u.sock = 0
|
||||
|
||||
windows.PostQueuedCompletionStatus(u.rx.iocp, 0, 0, nil)
|
||||
windows.PostQueuedCompletionStatus(u.tx.iocp, 0, 0, nil)
|
||||
|
||||
u.rx.CloseAndZero()
|
||||
u.tx.CloseAndZero()
|
||||
if u.sock != 0 {
|
||||
windows.CloseHandle(u.sock)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (ring *ringBuffer) Push() *ringPacket {
|
||||
|
||||
Reference in New Issue
Block a user