mirror of
https://github.com/slackhq/nebula.git
synced 2026-09-30 05:36:37 +02:00
Add nebula ctl, a local socket for the debug commands
Every diagnostic command nebula has was reachable through exactly one door:
the built-in ssh debug server. That server is off by default, and turning it
on means generating a host key, writing an sshd block with authorized public
keys, and SIGHUPing the daemon. That is a lot of ceremony to answer "what
version is this node running".
Nebula now serves the same commands over a local unix socket, enabled by
default, and `nebula ctl <command>` runs them. The socket lives in a 0700
directory so filesystem permissions are the access control; no keys, nothing
on the network. Failing to create it is logged and never blocks startup.
The command registry was already transport neutral, so this is mostly new
transport rather than new commands:
- diag/ holds the registry, dispatch, writer and wire protocol, moved out
of sshd because none of it was ever about ssh. sshd and ctl.go dispatch
against one shared registry.
- commands.go holds every command implementation, moved out of ssh.go
(which was 85% not ssh) and renamed off the ssh prefix. Adding a command
there makes it available over both transports.
- ssh.go keeps only host keys, authorized users, and the listen address.
- ctl.go supervises the socket, following the statsServer lifecycle shape.
The wire protocol frames the response rather than terminating it, because
print-cert -raw and list-hostmap -json both emit arbitrary bytes that no
sentinel could safely delimit. argv travels as a list so quoting survives.
Exit statuses are real: 0, 2 for usage, 127 for an unknown command.
Two things fall out. The ssh console now reports a real exit status instead
of a hardcoded zero, so `ssh host list-hostmap` is scriptable too. And eight
command callbacks that silently returned nil on a flags type mismatch now
report it, which the exit status makes visible.
Windows is a stub returning a clear "not supported" until it gets a named
pipe with a security descriptor; iOS and Android are never enabled, having no
daemon for a CLI to attach to.
Breaking for embedders of the sshd package: NewSSHServer takes a
*diag.Registry, SSHServer.RegisterCommand is gone in favor of registering on
that registry, and the command types live in diag rather than sshd.
Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014fya5fTXGiwX72FUmoL9y3
This commit is contained in:
co-authored by
Claude Opus 5
parent
89178f45ba
commit
14a4d87faf
@@ -7,6 +7,28 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
## [Unreleased]
|
## [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
|
## [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.
|
See the [v1.11.1](https://github.com/slackhq/nebula/milestone/30?closed=1) milestone for a complete list of changes.
|
||||||
|
|||||||
@@ -0,0 +1,131 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"flag"
|
||||||
|
"fmt"
|
||||||
|
"io/fs"
|
||||||
|
"log/slog"
|
||||||
|
"os"
|
||||||
|
"syscall"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/diag"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ctlMain implements `nebula ctl <command> [args...]`, which runs a debug command against the
|
||||||
|
// nebula already running on this host. Everything after the command name is forwarded to that
|
||||||
|
// nebula verbatim and parsed there by the same flag sets the ssh console uses, so this side
|
||||||
|
// deliberately understands as little as possible about it.
|
||||||
|
//
|
||||||
|
// Returns the process exit status.
|
||||||
|
func ctlMain(argv []string) int {
|
||||||
|
fl := flag.NewFlagSet("nebula ctl", flag.ContinueOnError)
|
||||||
|
fl.Usage = func() {
|
||||||
|
out := fl.Output()
|
||||||
|
fmt.Fprintf(out, "Usage: nebula ctl [-config path] [-socket path] <command> [arguments]\n\n")
|
||||||
|
fmt.Fprintf(out, "Runs a debug command against the running nebula on this host, over its local\n")
|
||||||
|
fmt.Fprintf(out, "control socket. Run `nebula ctl` with no command for the list of commands.\n\n")
|
||||||
|
fl.PrintDefaults()
|
||||||
|
}
|
||||||
|
|
||||||
|
socket := fl.String("socket", "", "Path to the control socket. Overrides ctl.socket from the config")
|
||||||
|
configPath := fl.String("config", "", "Path to the nebula config, read only to find ctl.socket")
|
||||||
|
|
||||||
|
// The flag package stops at the first non-flag argument, which is exactly the behaviour
|
||||||
|
// wanted here: `nebula ctl -socket /x list-hostmap -json` consumes -socket, stops at
|
||||||
|
// list-hostmap, and leaves the rest untouched for the daemon to parse.
|
||||||
|
if err := fl.Parse(argv); err != nil {
|
||||||
|
// -h is a request, not a failure.
|
||||||
|
if errors.Is(err, flag.ErrHelp) {
|
||||||
|
return diag.StatusOK
|
||||||
|
}
|
||||||
|
return diag.StatusUsage
|
||||||
|
}
|
||||||
|
|
||||||
|
path := *socket
|
||||||
|
if path == "" {
|
||||||
|
path = ctlSocketPath(*configPath)
|
||||||
|
}
|
||||||
|
|
||||||
|
if path == "" {
|
||||||
|
fmt.Fprintln(os.Stderr, "nebula ctl: no control socket path is known for this platform, set ctl.socket in the config")
|
||||||
|
return diag.StatusError
|
||||||
|
}
|
||||||
|
|
||||||
|
client, err := diag.Dial(path)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Fprintln(os.Stderr, ctlDialError(path, err))
|
||||||
|
return diag.StatusError
|
||||||
|
}
|
||||||
|
defer client.Close()
|
||||||
|
|
||||||
|
args := fl.Args()
|
||||||
|
status, err := client.Run(args, os.Stdout)
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, diag.ErrTruncated) {
|
||||||
|
fmt.Fprintf(os.Stderr, "nebula ctl: nebula closed the connection before %s finished\n", ctlCommandName(args))
|
||||||
|
return diag.StatusError
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintf(os.Stderr, "nebula ctl: %s\n", err)
|
||||||
|
if status == diag.StatusOK {
|
||||||
|
return diag.StatusError
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return status
|
||||||
|
}
|
||||||
|
|
||||||
|
// ctlSocketPath finds the socket to talk to. The platform default is the primary mechanism;
|
||||||
|
// reading the config is the refinement for someone who moved the socket. It is best effort by
|
||||||
|
// design, because config.DefaultPath resolves next to the nebula binary and a packaged install
|
||||||
|
// keeps its config somewhere else entirely, so a config we cannot find is the normal case
|
||||||
|
// rather than a failure.
|
||||||
|
func ctlSocketPath(configPath string) string {
|
||||||
|
if configPath == "" {
|
||||||
|
p, err := config.DefaultPath()
|
||||||
|
if err != nil {
|
||||||
|
return diag.DefaultSocketPath()
|
||||||
|
}
|
||||||
|
configPath = p
|
||||||
|
}
|
||||||
|
|
||||||
|
c := config.NewC(slog.New(slog.DiscardHandler))
|
||||||
|
if err := c.Load(configPath); err != nil {
|
||||||
|
return diag.DefaultSocketPath()
|
||||||
|
}
|
||||||
|
|
||||||
|
return c.GetString("ctl.socket", diag.DefaultSocketPath())
|
||||||
|
}
|
||||||
|
|
||||||
|
// ctlDialError turns a connect failure into something an operator can act on. These messages
|
||||||
|
// are the entire user experience when things are not working, so they name the path and say
|
||||||
|
// what to check.
|
||||||
|
func ctlDialError(path string, err error) string {
|
||||||
|
switch {
|
||||||
|
case errors.Is(err, diag.ErrNotSupported):
|
||||||
|
return "nebula ctl is not supported on this platform yet"
|
||||||
|
|
||||||
|
case errors.Is(err, fs.ErrNotExist):
|
||||||
|
return fmt.Sprintf("nebula ctl: no control socket at %s. Is nebula running? Is ctl.enabled set to false, or ctl.socket set to another path?", path)
|
||||||
|
|
||||||
|
case errors.Is(err, syscall.ECONNREFUSED):
|
||||||
|
return fmt.Sprintf("nebula ctl: found a stale socket at %s, nebula is not listening on it", path)
|
||||||
|
|
||||||
|
case errors.Is(err, fs.ErrPermission):
|
||||||
|
return fmt.Sprintf("nebula ctl: permission denied opening %s. nebula ctl must run as the user nebula runs as, usually root", path)
|
||||||
|
|
||||||
|
default:
|
||||||
|
return fmt.Sprintf("nebula ctl: %s", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ctlCommandName names the command for an error message, for the case where there isn't one.
|
||||||
|
func ctlCommandName(args []string) string {
|
||||||
|
if len(args) == 0 {
|
||||||
|
return "the command"
|
||||||
|
}
|
||||||
|
|
||||||
|
return args[0]
|
||||||
|
}
|
||||||
@@ -0,0 +1,70 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"io/fs"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"syscall"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/diag"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The daemon parses the command's own flags, so this side must consume its own and forward
|
||||||
|
// everything from the command name onwards untouched.
|
||||||
|
func TestCtlSocketPath(t *testing.T) {
|
||||||
|
t.Run("a config naming a socket is used", func(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
path := filepath.Join(dir, "config.yml")
|
||||||
|
require.NoError(t, os.WriteFile(path, []byte("ctl:\n socket: /run/somewhere/ctl.sock\n"), 0600))
|
||||||
|
|
||||||
|
assert.Equal(t, "/run/somewhere/ctl.sock", ctlSocketPath(path))
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a config without a ctl block falls back to the platform default", func(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
path := filepath.Join(dir, "config.yml")
|
||||||
|
require.NoError(t, os.WriteFile(path, []byte("pki:\n ca: /dev/null\n"), 0600))
|
||||||
|
|
||||||
|
assert.Equal(t, diag.DefaultSocketPath(), ctlSocketPath(path))
|
||||||
|
})
|
||||||
|
|
||||||
|
// A packaged install keeps its config somewhere config.DefaultPath will never look, so a
|
||||||
|
// config we cannot read is the ordinary case and must not be fatal.
|
||||||
|
t.Run("an unreadable config falls back to the platform default", func(t *testing.T) {
|
||||||
|
assert.Equal(t, diag.DefaultSocketPath(), ctlSocketPath(filepath.Join(t.TempDir(), "nope.yml")))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCtlDialError(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
err error
|
||||||
|
wants string
|
||||||
|
}{
|
||||||
|
{"missing socket names the path and what to check", fs.ErrNotExist, "no control socket at /x/ctl.sock. Is nebula running?"},
|
||||||
|
{"a stale socket is called stale", syscall.ECONNREFUSED, "found a stale socket at /x/ctl.sock"},
|
||||||
|
{"permission denied suggests the right user", fs.ErrPermission, "must run as the user nebula runs as"},
|
||||||
|
{"an unsupported platform says so", diag.ErrNotSupported, "not supported on this platform"},
|
||||||
|
{"anything else is reported verbatim", errors.New("something else"), "something else"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
assert.Contains(t, ctlDialError("/x/ctl.sock", tt.err), tt.wants)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("a wrapped syscall error is still recognised", func(t *testing.T) {
|
||||||
|
err := &os.SyscallError{Syscall: "connect", Err: syscall.ECONNREFUSED}
|
||||||
|
assert.Contains(t, ctlDialError("/x/ctl.sock", err), "stale socket")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCtlCommandName(t *testing.T) {
|
||||||
|
assert.Equal(t, "print-cert", ctlCommandName([]string{"print-cert", "-json"}))
|
||||||
|
assert.Equal(t, "the command", ctlCommandName(nil))
|
||||||
|
}
|
||||||
@@ -32,11 +32,26 @@ func init() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func main() {
|
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")
|
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")
|
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")
|
printVersion := flag.Bool("version", false, "Print version")
|
||||||
printUsage := flag.Bool("help", false, "Print command line usage")
|
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()
|
flag.Parse()
|
||||||
|
|
||||||
if *printVersion {
|
if *printVersion {
|
||||||
|
|||||||
+922
@@ -0,0 +1,922 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
// The commands nebula exposes for debugging and administration. They are transport neutral:
|
||||||
|
// the ssh console in ssh.go and the `nebula ctl` socket in ctl.go both dispatch against the
|
||||||
|
// registry attachCommands fills in, and a command cannot tell which one invoked it. Adding a
|
||||||
|
// command here makes it available over both.
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"flag"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"maps"
|
||||||
|
"net/netip"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"runtime"
|
||||||
|
"runtime/pprof"
|
||||||
|
"sort"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/diag"
|
||||||
|
"github.com/slackhq/nebula/header"
|
||||||
|
"github.com/slackhq/nebula/logging"
|
||||||
|
)
|
||||||
|
|
||||||
|
type listHostMapFlags struct {
|
||||||
|
Json bool
|
||||||
|
Pretty bool
|
||||||
|
ByIndex bool
|
||||||
|
}
|
||||||
|
|
||||||
|
type printCertFlags struct {
|
||||||
|
Json bool
|
||||||
|
Pretty bool
|
||||||
|
Raw bool
|
||||||
|
}
|
||||||
|
|
||||||
|
type printTunnelFlags struct {
|
||||||
|
Pretty bool
|
||||||
|
}
|
||||||
|
|
||||||
|
type changeRemoteFlags struct {
|
||||||
|
Address string
|
||||||
|
}
|
||||||
|
|
||||||
|
type closeTunnelFlags struct {
|
||||||
|
LocalOnly bool
|
||||||
|
}
|
||||||
|
|
||||||
|
type createTunnelFlags struct {
|
||||||
|
Address string
|
||||||
|
}
|
||||||
|
|
||||||
|
type deviceInfoFlags struct {
|
||||||
|
Json bool
|
||||||
|
Pretty bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func attachCommands(l *slog.Logger, c *config.C, reg *diag.Registry, f *Interface) {
|
||||||
|
// sandboxDir defaults to a dir in temp. The intention is that end user will
|
||||||
|
// create this dir as needed. Overriding this config value to "" allows
|
||||||
|
// writing to anywhere in the system.
|
||||||
|
defaultDir := filepath.Join(os.TempDir(), "nebula-debug")
|
||||||
|
// The key is spelled for both transports now: the profile writers are reachable over
|
||||||
|
// `nebula ctl` as well, but sshd.sandbox_dir keeps working for anyone already setting it.
|
||||||
|
sandboxDir := c.GetString("ctl.sandbox_dir", c.GetString("sshd.sandbox_dir", defaultDir))
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "list-hostmap",
|
||||||
|
ShortDescription: "List all known previously connected hosts",
|
||||||
|
Flags: func() (*flag.FlagSet, any) {
|
||||||
|
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||||
|
s := listHostMapFlags{}
|
||||||
|
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
|
||||||
|
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
|
||||||
|
fl.BoolVar(&s.ByIndex, "by-index", false, "gets all hosts in the hostmap from the index table")
|
||||||
|
return fl, &s
|
||||||
|
},
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdListHostMap(f.hostMap, fs, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "list-pending-hostmap",
|
||||||
|
ShortDescription: "List all handshaking hosts",
|
||||||
|
Flags: func() (*flag.FlagSet, any) {
|
||||||
|
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||||
|
s := listHostMapFlags{}
|
||||||
|
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
|
||||||
|
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
|
||||||
|
fl.BoolVar(&s.ByIndex, "by-index", false, "gets all hosts in the hostmap from the index table")
|
||||||
|
return fl, &s
|
||||||
|
},
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdListHostMap(f.handshakeManager, fs, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "list-lighthouse-addrmap",
|
||||||
|
ShortDescription: "List all lighthouse map entries",
|
||||||
|
Flags: func() (*flag.FlagSet, any) {
|
||||||
|
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||||
|
s := listHostMapFlags{}
|
||||||
|
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
|
||||||
|
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
|
||||||
|
return fl, &s
|
||||||
|
},
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdListLighthouseMap(f.lightHouse, fs, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "reload",
|
||||||
|
ShortDescription: "Reloads configuration from disk, same as sending HUP to the process",
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdReload(c, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "start-cpu-profile",
|
||||||
|
ShortDescription: "Starts a cpu profile and write output to the provided file, ex: `cpu-profile.pb.gz`",
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdStartCpuProfile(sandboxDir, fs, a, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "stop-cpu-profile",
|
||||||
|
ShortDescription: "Stops a cpu profile and writes output to the previously provided file",
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
pprof.StopCPUProfile()
|
||||||
|
return w.WriteLine("If a CPU profile was running it is now stopped")
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "save-heap-profile",
|
||||||
|
ShortDescription: "Saves a heap profile to the provided path, ex: `heap-profile.pb.gz`",
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdGetHeapProfile(sandboxDir, fs, a, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "mutex-profile-fraction",
|
||||||
|
ShortDescription: "Gets or sets runtime.SetMutexProfileFraction",
|
||||||
|
Callback: cmdMutexProfileFraction,
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "save-mutex-profile",
|
||||||
|
ShortDescription: "Saves a mutex profile to the provided path, ex: `mutex-profile.pb.gz`",
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdGetMutexProfile(sandboxDir, fs, a, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "log-level",
|
||||||
|
ShortDescription: "Gets or sets the current log level",
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdLogLevel(l, fs, a, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "log-format",
|
||||||
|
ShortDescription: "Gets or sets the current log format",
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdLogFormat(l, fs, a, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "version",
|
||||||
|
ShortDescription: "Prints the currently running version of nebula",
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdVersion(f, fs, a, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "device-info",
|
||||||
|
ShortDescription: "Prints information about the network device.",
|
||||||
|
Flags: func() (*flag.FlagSet, any) {
|
||||||
|
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||||
|
s := deviceInfoFlags{}
|
||||||
|
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
|
||||||
|
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
|
||||||
|
return fl, &s
|
||||||
|
},
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdDeviceInfo(f, fs, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "print-cert",
|
||||||
|
ShortDescription: "Prints the current certificate being used or the certificate for the provided vpn addr",
|
||||||
|
Flags: func() (*flag.FlagSet, any) {
|
||||||
|
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||||
|
s := printCertFlags{}
|
||||||
|
fl.BoolVar(&s.Json, "json", false, "outputs as json")
|
||||||
|
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
|
||||||
|
fl.BoolVar(&s.Raw, "raw", false, "raw prints the PEM encoded certificate, not compatible with -json or -pretty")
|
||||||
|
return fl, &s
|
||||||
|
},
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdPrintCert(f, fs, a, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "print-tunnel",
|
||||||
|
ShortDescription: "Prints json details about a tunnel for the provided vpn addr",
|
||||||
|
Flags: func() (*flag.FlagSet, any) {
|
||||||
|
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||||
|
s := printTunnelFlags{}
|
||||||
|
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json")
|
||||||
|
return fl, &s
|
||||||
|
},
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdPrintTunnel(f, fs, a, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "print-relays",
|
||||||
|
ShortDescription: "Prints json details about all relay info",
|
||||||
|
Flags: func() (*flag.FlagSet, any) {
|
||||||
|
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||||
|
s := printTunnelFlags{}
|
||||||
|
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json")
|
||||||
|
return fl, &s
|
||||||
|
},
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdPrintRelays(f, fs, a, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "change-remote",
|
||||||
|
ShortDescription: "Changes the remote address used in the tunnel for the provided vpn addr",
|
||||||
|
Flags: func() (*flag.FlagSet, any) {
|
||||||
|
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||||
|
s := changeRemoteFlags{}
|
||||||
|
fl.StringVar(&s.Address, "address", "", "The new remote address, ip:port")
|
||||||
|
return fl, &s
|
||||||
|
},
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdChangeRemote(f, fs, a, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "close-tunnel",
|
||||||
|
ShortDescription: "Closes a tunnel for the provided vpn addr",
|
||||||
|
Flags: func() (*flag.FlagSet, any) {
|
||||||
|
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||||
|
s := closeTunnelFlags{}
|
||||||
|
fl.BoolVar(&s.LocalOnly, "local-only", false, "Disables notifying the remote that the tunnel is shutting down")
|
||||||
|
return fl, &s
|
||||||
|
},
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdCloseTunnel(f, fs, a, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "create-tunnel",
|
||||||
|
ShortDescription: "Creates a tunnel for the provided vpn address",
|
||||||
|
Help: "The lighthouses will be queried for real addresses but you can provide one as well.",
|
||||||
|
Flags: func() (*flag.FlagSet, any) {
|
||||||
|
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||||
|
s := createTunnelFlags{}
|
||||||
|
fl.StringVar(&s.Address, "address", "", "Optionally provide a real remote address, ip:port ")
|
||||||
|
return fl, &s
|
||||||
|
},
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdCreateTunnel(f, fs, a, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
reg.RegisterCommand(&diag.Command{
|
||||||
|
Name: "query-lighthouse",
|
||||||
|
ShortDescription: "Query the lighthouses for the provided vpn address",
|
||||||
|
Help: "This command is asynchronous. Only currently known udp addresses will be printed.",
|
||||||
|
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return cmdQueryLighthouse(f, fs, a, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdListHostMap(hl controlHostLister, a any, w diag.StringWriter) error {
|
||||||
|
fs, ok := a.(*listHostMapFlags)
|
||||||
|
if !ok {
|
||||||
|
return fmt.Errorf("internal error: expected flags to be listHostMapFlags but was %+v", a)
|
||||||
|
}
|
||||||
|
|
||||||
|
var hm []ControlHostInfo
|
||||||
|
if fs.ByIndex {
|
||||||
|
hm = listHostMapIndexes(hl)
|
||||||
|
} else {
|
||||||
|
hm = listHostMapHosts(hl)
|
||||||
|
}
|
||||||
|
|
||||||
|
sort.Slice(hm, func(i, j int) bool {
|
||||||
|
return hm[i].VpnAddrs[0].Compare(hm[j].VpnAddrs[0]) < 0
|
||||||
|
})
|
||||||
|
|
||||||
|
if fs.Json || fs.Pretty {
|
||||||
|
js := json.NewEncoder(w.GetWriter())
|
||||||
|
if fs.Pretty {
|
||||||
|
js.SetIndent("", " ")
|
||||||
|
}
|
||||||
|
|
||||||
|
err := js.Encode(hm)
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
} else {
|
||||||
|
for _, v := range hm {
|
||||||
|
err := w.WriteLine(fmt.Sprintf("%s: %s", v.VpnAddrs, v.RemoteAddrs))
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdListLighthouseMap(lightHouse *LightHouse, a any, w diag.StringWriter) error {
|
||||||
|
fs, ok := a.(*listHostMapFlags)
|
||||||
|
if !ok {
|
||||||
|
return fmt.Errorf("internal error: expected flags to be listHostMapFlags but was %+v", a)
|
||||||
|
}
|
||||||
|
|
||||||
|
type lighthouseInfo struct {
|
||||||
|
VpnAddr string `json:"vpnAddr"`
|
||||||
|
Addrs *CacheMap `json:"addrs"`
|
||||||
|
}
|
||||||
|
|
||||||
|
lightHouse.RLock()
|
||||||
|
addrMap := make([]lighthouseInfo, len(lightHouse.addrMap))
|
||||||
|
x := 0
|
||||||
|
for k, v := range lightHouse.addrMap {
|
||||||
|
addrMap[x] = lighthouseInfo{
|
||||||
|
VpnAddr: k.String(),
|
||||||
|
Addrs: v.CopyCache(),
|
||||||
|
}
|
||||||
|
x++
|
||||||
|
}
|
||||||
|
lightHouse.RUnlock()
|
||||||
|
|
||||||
|
sort.Slice(addrMap, func(i, j int) bool {
|
||||||
|
return strings.Compare(addrMap[i].VpnAddr, addrMap[j].VpnAddr) < 0
|
||||||
|
})
|
||||||
|
|
||||||
|
if fs.Json || fs.Pretty {
|
||||||
|
js := json.NewEncoder(w.GetWriter())
|
||||||
|
if fs.Pretty {
|
||||||
|
js.SetIndent("", " ")
|
||||||
|
}
|
||||||
|
|
||||||
|
err := js.Encode(addrMap)
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
} else {
|
||||||
|
for _, v := range addrMap {
|
||||||
|
b, err := json.Marshal(v.Addrs)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
err = w.WriteLine(fmt.Sprintf("%s: %s", v.VpnAddr, string(b)))
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// sanitizeFilePath validates that the given file path is within the sandbox directory.
|
||||||
|
// If sandboxDir is empty, the path is returned as-is for backwards compatibility.
|
||||||
|
func sanitizeFilePath(sandboxDir, filePath string) (string, error) {
|
||||||
|
if sandboxDir == "" {
|
||||||
|
return filePath, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clean and resolve the path relative to the sandbox directory
|
||||||
|
if !filepath.IsAbs(filePath) {
|
||||||
|
filePath = filepath.Join(sandboxDir, filePath)
|
||||||
|
}
|
||||||
|
cleaned := filepath.Clean(filePath)
|
||||||
|
|
||||||
|
// Ensure the resolved path is within the sandbox directory
|
||||||
|
cleanedSandbox := filepath.Clean(sandboxDir)
|
||||||
|
if cleaned == cleanedSandbox {
|
||||||
|
return "", fmt.Errorf("path %q resolves to the sandbox directory itself %q", filePath, sandboxDir)
|
||||||
|
}
|
||||||
|
if !strings.HasPrefix(cleaned, cleanedSandbox+string(filepath.Separator)) {
|
||||||
|
return "", fmt.Errorf("path %q is outside the sandbox directory %q", filePath, sandboxDir)
|
||||||
|
}
|
||||||
|
|
||||||
|
return cleaned, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdStartCpuProfile(sandboxDir string, fs any, a []string, w diag.StringWriter) error {
|
||||||
|
if len(a) == 0 {
|
||||||
|
err := w.WriteLine("No path to write profile provided")
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
filePath, err := sanitizeFilePath(sandboxDir, a[0])
|
||||||
|
if err != nil {
|
||||||
|
return w.WriteLine(err.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
file, err := os.Create(filePath)
|
||||||
|
if err != nil {
|
||||||
|
err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
err = pprof.StartCPUProfile(file)
|
||||||
|
if err != nil {
|
||||||
|
err = w.WriteLine(fmt.Sprintf("Unable to start cpu profile: %s", err))
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
err = w.WriteLine(fmt.Sprintf("Started cpu profile, issue stop-cpu-profile to write the output to %s", a))
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdVersion(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||||
|
return w.WriteLine(fmt.Sprintf("%s", ifce.version))
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdQueryLighthouse(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||||
|
if len(a) == 0 {
|
||||||
|
return w.WriteLine("No vpn address was provided")
|
||||||
|
}
|
||||||
|
|
||||||
|
vpnAddr, err := netip.ParseAddr(a[0])
|
||||||
|
if err != nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
if !vpnAddr.IsValid() {
|
||||||
|
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
var cm *CacheMap
|
||||||
|
rl := ifce.lightHouse.Query(vpnAddr)
|
||||||
|
if rl != nil {
|
||||||
|
cm = rl.CopyCache()
|
||||||
|
}
|
||||||
|
return json.NewEncoder(w.GetWriter()).Encode(cm)
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdCloseTunnel(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||||
|
flags, ok := fs.(*closeTunnelFlags)
|
||||||
|
if !ok {
|
||||||
|
return fmt.Errorf("internal error: expected flags to be closeTunnelFlags but was %+v", fs)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(a) == 0 {
|
||||||
|
return w.WriteLine("No vpn address was provided")
|
||||||
|
}
|
||||||
|
|
||||||
|
vpnAddr, err := netip.ParseAddr(a[0])
|
||||||
|
if err != nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
if !vpnAddr.IsValid() {
|
||||||
|
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||||
|
if hostInfo == nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn address: %v", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
if !flags.LocalOnly {
|
||||||
|
ifce.send(
|
||||||
|
header.CloseTunnel,
|
||||||
|
0,
|
||||||
|
hostInfo.ConnectionState,
|
||||||
|
hostInfo,
|
||||||
|
[]byte{},
|
||||||
|
make([]byte, 12, 12),
|
||||||
|
make([]byte, mtu),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
ifce.closeTunnel(hostInfo)
|
||||||
|
return w.WriteLine("Closed")
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdCreateTunnel(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||||
|
flags, ok := fs.(*createTunnelFlags)
|
||||||
|
if !ok {
|
||||||
|
return fmt.Errorf("internal error: expected flags to be createTunnelFlags but was %+v", fs)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(a) == 0 {
|
||||||
|
return w.WriteLine("No vpn address was provided")
|
||||||
|
}
|
||||||
|
|
||||||
|
vpnAddr, err := netip.ParseAddr(a[0])
|
||||||
|
if err != nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
if !vpnAddr.IsValid() {
|
||||||
|
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||||
|
if hostInfo != nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("Tunnel already exists"))
|
||||||
|
}
|
||||||
|
|
||||||
|
hostInfo = ifce.handshakeManager.QueryVpnAddr(vpnAddr)
|
||||||
|
if hostInfo != nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("Tunnel already handshaking"))
|
||||||
|
}
|
||||||
|
|
||||||
|
var addr netip.AddrPort
|
||||||
|
if flags.Address != "" {
|
||||||
|
addr, err = netip.ParseAddrPort(flags.Address)
|
||||||
|
if err != nil {
|
||||||
|
return w.WriteLine("Address could not be parsed")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
hostInfo = ifce.handshakeManager.StartHandshake(vpnAddr, nil)
|
||||||
|
if addr.IsValid() {
|
||||||
|
hostInfo.SetRemote(addr)
|
||||||
|
}
|
||||||
|
|
||||||
|
return w.WriteLine("Created")
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdChangeRemote(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||||
|
flags, ok := fs.(*changeRemoteFlags)
|
||||||
|
if !ok {
|
||||||
|
return fmt.Errorf("internal error: expected flags to be changeRemoteFlags but was %+v", fs)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(a) == 0 {
|
||||||
|
return w.WriteLine("No vpn address was provided")
|
||||||
|
}
|
||||||
|
|
||||||
|
if flags.Address == "" {
|
||||||
|
return w.WriteLine("No address was provided")
|
||||||
|
}
|
||||||
|
|
||||||
|
addr, err := netip.ParseAddrPort(flags.Address)
|
||||||
|
if err != nil {
|
||||||
|
return w.WriteLine("Address could not be parsed")
|
||||||
|
}
|
||||||
|
|
||||||
|
vpnAddr, err := netip.ParseAddr(a[0])
|
||||||
|
if err != nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
if !vpnAddr.IsValid() {
|
||||||
|
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||||
|
if hostInfo == nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn address: %v", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
hostInfo.SetRemote(addr)
|
||||||
|
return w.WriteLine("Changed")
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdGetHeapProfile(sandboxDir string, fs any, a []string, w diag.StringWriter) error {
|
||||||
|
if len(a) == 0 {
|
||||||
|
return w.WriteLine("No path to write profile provided")
|
||||||
|
}
|
||||||
|
|
||||||
|
filePath, err := sanitizeFilePath(sandboxDir, a[0])
|
||||||
|
if err != nil {
|
||||||
|
return w.WriteLine(err.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
file, err := os.Create(filePath)
|
||||||
|
if err != nil {
|
||||||
|
err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
err = pprof.WriteHeapProfile(file)
|
||||||
|
if err != nil {
|
||||||
|
err = w.WriteLine(fmt.Sprintf("Unable to write profile: %s", err))
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
err = w.WriteLine(fmt.Sprintf("Mem profile created at %s", a))
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdMutexProfileFraction(fs any, a []string, w diag.StringWriter) error {
|
||||||
|
if len(a) == 0 {
|
||||||
|
rate := runtime.SetMutexProfileFraction(-1)
|
||||||
|
return w.WriteLine(fmt.Sprintf("Current value: %d", rate))
|
||||||
|
}
|
||||||
|
|
||||||
|
newRate, err := strconv.Atoi(a[0])
|
||||||
|
if err != nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("Invalid argument: %s", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
oldRate := runtime.SetMutexProfileFraction(newRate)
|
||||||
|
return w.WriteLine(fmt.Sprintf("New value: %d. Old value: %d", newRate, oldRate))
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdGetMutexProfile(sandboxDir string, fs any, a []string, w diag.StringWriter) error {
|
||||||
|
if len(a) == 0 {
|
||||||
|
return w.WriteLine("No path to write profile provided")
|
||||||
|
}
|
||||||
|
|
||||||
|
filePath, err := sanitizeFilePath(sandboxDir, a[0])
|
||||||
|
if err != nil {
|
||||||
|
return w.WriteLine(err.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
file, err := os.Create(filePath)
|
||||||
|
if err != nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
|
||||||
|
}
|
||||||
|
defer file.Close()
|
||||||
|
|
||||||
|
mutexProfile := pprof.Lookup("mutex")
|
||||||
|
if mutexProfile == nil {
|
||||||
|
return w.WriteLine("Unable to get pprof.Lookup(\"mutex\")")
|
||||||
|
}
|
||||||
|
|
||||||
|
err = mutexProfile.WriteTo(file, 0)
|
||||||
|
if err != nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("Unable to write profile: %s", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
return w.WriteLine(fmt.Sprintf("Mutex profile created at %s", a))
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdLogLevel(l *slog.Logger, fs any, a []string, w diag.StringWriter) error {
|
||||||
|
ctrl, ok := l.Handler().(interface {
|
||||||
|
GetLevel() slog.Level
|
||||||
|
SetLevel(slog.Level)
|
||||||
|
})
|
||||||
|
if !ok {
|
||||||
|
return w.WriteLine("Log level is not reconfigurable on this logger")
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(a) == 0 {
|
||||||
|
return w.WriteLine(fmt.Sprintf("Log level is: %s", logging.LevelName(ctrl.GetLevel())))
|
||||||
|
}
|
||||||
|
|
||||||
|
level, err := logging.ParseLevel(strings.ToLower(a[0]))
|
||||||
|
if err != nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("Unknown log level %s. Possible log levels: trace, debug, info, warn, error", a))
|
||||||
|
}
|
||||||
|
|
||||||
|
ctrl.SetLevel(level)
|
||||||
|
return w.WriteLine(fmt.Sprintf("Log level is: %s", logging.LevelName(ctrl.GetLevel())))
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdLogFormat(l *slog.Logger, fs any, a []string, w diag.StringWriter) error {
|
||||||
|
ctrl, ok := l.Handler().(interface {
|
||||||
|
GetFormat() string
|
||||||
|
SetFormat(string) error
|
||||||
|
})
|
||||||
|
if !ok {
|
||||||
|
return w.WriteLine("Log format is not reconfigurable on this logger")
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(a) == 0 {
|
||||||
|
return w.WriteLine(fmt.Sprintf("Log format is: %s", ctrl.GetFormat()))
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := ctrl.SetFormat(strings.ToLower(a[0])); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return w.WriteLine(fmt.Sprintf("Log format is: %s", ctrl.GetFormat()))
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdPrintCert(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||||
|
args, ok := fs.(*printCertFlags)
|
||||||
|
if !ok {
|
||||||
|
return fmt.Errorf("internal error: expected flags to be printCertFlags but was %+v", fs)
|
||||||
|
}
|
||||||
|
|
||||||
|
cert := ifce.pki.getCertState().GetDefaultCertificate()
|
||||||
|
if len(a) > 0 {
|
||||||
|
vpnAddr, err := netip.ParseAddr(a[0])
|
||||||
|
if err != nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
if !vpnAddr.IsValid() {
|
||||||
|
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||||
|
if hostInfo == nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn addr: %v", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
cert = hostInfo.GetCert().Certificate
|
||||||
|
}
|
||||||
|
|
||||||
|
if args.Json || args.Pretty {
|
||||||
|
b, err := cert.MarshalJSON()
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if args.Pretty {
|
||||||
|
buf := new(bytes.Buffer)
|
||||||
|
err := json.Indent(buf, b, "", " ")
|
||||||
|
b = buf.Bytes()
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return w.WriteBytes(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
if args.Raw {
|
||||||
|
b, err := cert.MarshalPEM()
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return w.WriteBytes(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
return w.WriteLine(cert.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdPrintRelays(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||||
|
args, ok := fs.(*printTunnelFlags)
|
||||||
|
if !ok {
|
||||||
|
return fmt.Errorf("internal error: expected flags to be printTunnelFlags but was %+v", fs)
|
||||||
|
}
|
||||||
|
|
||||||
|
relays := map[uint32]*HostInfo{}
|
||||||
|
ifce.hostMap.Lock()
|
||||||
|
maps.Copy(relays, ifce.hostMap.Relays)
|
||||||
|
ifce.hostMap.Unlock()
|
||||||
|
|
||||||
|
type RelayFor struct {
|
||||||
|
Error error
|
||||||
|
Type string
|
||||||
|
State string
|
||||||
|
PeerAddr netip.Addr
|
||||||
|
LocalIndex uint32
|
||||||
|
RemoteIndex uint32
|
||||||
|
RelayedThrough []netip.Addr
|
||||||
|
}
|
||||||
|
|
||||||
|
type RelayOutput struct {
|
||||||
|
NebulaAddr netip.Addr
|
||||||
|
RelayForAddrs []RelayFor
|
||||||
|
}
|
||||||
|
|
||||||
|
type CmdOutput struct {
|
||||||
|
Relays []*RelayOutput
|
||||||
|
}
|
||||||
|
|
||||||
|
co := CmdOutput{}
|
||||||
|
|
||||||
|
enc := json.NewEncoder(w.GetWriter())
|
||||||
|
|
||||||
|
if args.Pretty {
|
||||||
|
enc.SetIndent("", " ")
|
||||||
|
}
|
||||||
|
|
||||||
|
for k, v := range relays {
|
||||||
|
ro := RelayOutput{NebulaAddr: v.vpnAddrs[0]}
|
||||||
|
co.Relays = append(co.Relays, &ro)
|
||||||
|
relayHI := ifce.hostMap.QueryVpnAddr(v.vpnAddrs[0])
|
||||||
|
if relayHI == nil {
|
||||||
|
ro.RelayForAddrs = append(ro.RelayForAddrs, RelayFor{Error: errors.New("could not find hostinfo")})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for _, vpnAddr := range relayHI.relayState.CopyRelayForIps() {
|
||||||
|
rf := RelayFor{Error: nil}
|
||||||
|
r, ok := relayHI.relayState.GetRelayForByAddr(vpnAddr)
|
||||||
|
if ok {
|
||||||
|
t := ""
|
||||||
|
switch r.Type {
|
||||||
|
case ForwardingType:
|
||||||
|
t = "forwarding"
|
||||||
|
case TerminalType:
|
||||||
|
t = "terminal"
|
||||||
|
default:
|
||||||
|
t = "unknown"
|
||||||
|
}
|
||||||
|
|
||||||
|
s := ""
|
||||||
|
switch r.State {
|
||||||
|
case Requested:
|
||||||
|
s = "requested"
|
||||||
|
case Established:
|
||||||
|
s = "established"
|
||||||
|
default:
|
||||||
|
s = "unknown"
|
||||||
|
}
|
||||||
|
|
||||||
|
rf.LocalIndex = r.LocalIndex
|
||||||
|
rf.RemoteIndex = r.RemoteIndex
|
||||||
|
rf.PeerAddr = r.PeerAddr
|
||||||
|
rf.Type = t
|
||||||
|
rf.State = s
|
||||||
|
if rf.LocalIndex != k {
|
||||||
|
rf.Error = fmt.Errorf("hostmap LocalIndex '%v' does not match RelayState LocalIndex", k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
relayedHI := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||||
|
if relayedHI != nil {
|
||||||
|
rf.RelayedThrough = append(rf.RelayedThrough, relayedHI.relayState.CopyRelayIps()...)
|
||||||
|
}
|
||||||
|
|
||||||
|
ro.RelayForAddrs = append(ro.RelayForAddrs, rf)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
err := enc.Encode(co)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdPrintTunnel(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||||
|
args, ok := fs.(*printTunnelFlags)
|
||||||
|
if !ok {
|
||||||
|
return fmt.Errorf("internal error: expected flags to be printTunnelFlags but was %+v", fs)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(a) == 0 {
|
||||||
|
return w.WriteLine("No vpn address was provided")
|
||||||
|
}
|
||||||
|
|
||||||
|
vpnAddr, err := netip.ParseAddr(a[0])
|
||||||
|
if err != nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
if !vpnAddr.IsValid() {
|
||||||
|
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||||
|
if hostInfo == nil {
|
||||||
|
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn addr: %v", a[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
enc := json.NewEncoder(w.GetWriter())
|
||||||
|
if args.Pretty {
|
||||||
|
enc.SetIndent("", " ")
|
||||||
|
}
|
||||||
|
|
||||||
|
return enc.Encode(copyHostInfo(hostInfo, ifce.hostMap.GetPreferredRanges()))
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdDeviceInfo(ifce *Interface, fs any, w diag.StringWriter) error {
|
||||||
|
|
||||||
|
data := struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Cidr []netip.Prefix `json:"cidr"`
|
||||||
|
}{
|
||||||
|
Name: ifce.inside.Name(),
|
||||||
|
Cidr: make([]netip.Prefix, len(ifce.inside.Networks())),
|
||||||
|
}
|
||||||
|
|
||||||
|
copy(data.Cidr, ifce.inside.Networks())
|
||||||
|
|
||||||
|
flags, ok := fs.(*deviceInfoFlags)
|
||||||
|
if !ok {
|
||||||
|
return fmt.Errorf("internal error: expected flags to be deviceInfoFlags but was %+v", fs)
|
||||||
|
}
|
||||||
|
|
||||||
|
if flags.Json || flags.Pretty {
|
||||||
|
js := json.NewEncoder(w.GetWriter())
|
||||||
|
if flags.Pretty {
|
||||||
|
js.SetIndent("", " ")
|
||||||
|
}
|
||||||
|
|
||||||
|
return js.Encode(data)
|
||||||
|
} else {
|
||||||
|
return w.WriteLine(fmt.Sprintf("name=%v cidr=%v", data.Name, data.Cidr))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdReload(c *config.C, w diag.StringWriter) error {
|
||||||
|
err := w.WriteLine("Reloading config")
|
||||||
|
c.ReloadConfig()
|
||||||
|
return err
|
||||||
|
}
|
||||||
@@ -0,0 +1,69 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"log/slog"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/diag"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// attachedCommands is every command nebula exposes. The ssh console and `nebula ctl` dispatch
|
||||||
|
// against this one set, so this list is the contract for both transports.
|
||||||
|
var attachedCommands = []string{
|
||||||
|
"change-remote",
|
||||||
|
"close-tunnel",
|
||||||
|
"create-tunnel",
|
||||||
|
"device-info",
|
||||||
|
"list-hostmap",
|
||||||
|
"list-lighthouse-addrmap",
|
||||||
|
"list-pending-hostmap",
|
||||||
|
"log-format",
|
||||||
|
"log-level",
|
||||||
|
"mutex-profile-fraction",
|
||||||
|
"print-cert",
|
||||||
|
"print-relays",
|
||||||
|
"print-tunnel",
|
||||||
|
"query-lighthouse",
|
||||||
|
"reload",
|
||||||
|
"save-heap-profile",
|
||||||
|
"save-mutex-profile",
|
||||||
|
"start-cpu-profile",
|
||||||
|
"stop-cpu-profile",
|
||||||
|
"version",
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAttachCommands(t *testing.T) {
|
||||||
|
l := slog.New(slog.DiscardHandler)
|
||||||
|
reg := diag.NewRegistry()
|
||||||
|
|
||||||
|
// The callbacks capture these but do not touch them until a command runs, and this test
|
||||||
|
// only registers and asks for help.
|
||||||
|
attachCommands(l, config.NewC(l), reg, &Interface{})
|
||||||
|
|
||||||
|
t.Run("every command is registered", func(t *testing.T) {
|
||||||
|
for _, name := range attachedCommands {
|
||||||
|
assert.Equal(t, []string{name}, reg.Match(name), "%s is not registered", name)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("help is available for every command", func(t *testing.T) {
|
||||||
|
for _, name := range attachedCommands {
|
||||||
|
buf := &bytes.Buffer{}
|
||||||
|
require.NoError(t, reg.DispatchArgs([]string{"help", name}, diag.NewWriter(buf)), name)
|
||||||
|
assert.Contains(t, buf.String(), name+" - ", name)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("the command list names them all", func(t *testing.T) {
|
||||||
|
buf := &bytes.Buffer{}
|
||||||
|
require.NoError(t, reg.DispatchArgs(nil, diag.NewWriter(buf)))
|
||||||
|
|
||||||
|
for _, name := range attachedCommands {
|
||||||
|
assert.Contains(t, buf.String(), name+" - ", name)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -50,6 +50,7 @@ type Control struct {
|
|||||||
ctx context.Context
|
ctx context.Context
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
sshStart func()
|
sshStart func()
|
||||||
|
ctlStart func()
|
||||||
statsStart func()
|
statsStart func()
|
||||||
dnsStart func()
|
dnsStart func()
|
||||||
lighthouseStart func()
|
lighthouseStart func()
|
||||||
@@ -99,6 +100,9 @@ func (c *Control) Start() error {
|
|||||||
if c.sshStart != nil {
|
if c.sshStart != nil {
|
||||||
go c.sshStart()
|
go c.sshStart()
|
||||||
}
|
}
|
||||||
|
if c.ctlStart != nil {
|
||||||
|
go c.ctlStart()
|
||||||
|
}
|
||||||
if c.statsStart != nil {
|
if c.statsStart != nil {
|
||||||
go c.statsStart()
|
go c.statsStart()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,234 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"path/filepath"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/diag"
|
||||||
|
"github.com/slackhq/nebula/util"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ctlConfig is the parsed form of the `ctl` config block. It is comparable so that a reload
|
||||||
|
// can tell "nothing changed" from "the socket moved" with ==.
|
||||||
|
type ctlConfig struct {
|
||||||
|
enabled bool
|
||||||
|
socket string
|
||||||
|
|
||||||
|
// explicit records that the operator named a socket path rather than taking the platform
|
||||||
|
// default. It only affects how loudly a failure to listen is reported: an unprivileged
|
||||||
|
// nebula that cannot create /run/nebula is a normal deployment, not a problem to shout
|
||||||
|
// about on every upgrade, but a path someone chose deliberately failing to bind is.
|
||||||
|
explicit bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// ctlServer owns the unix socket `nebula ctl` connects to. It exposes the same command
|
||||||
|
// registry the ssh console does, minus the ceremony of running an ssh server: the socket is
|
||||||
|
// local only and guarded by filesystem permissions, so it needs no keys.
|
||||||
|
//
|
||||||
|
// The lifecycle mirrors statsServer: the constructor wires the reload callback, reload
|
||||||
|
// records config and reconciles a running listener, Start builds and serves the runtime, and
|
||||||
|
// Stop tears it down.
|
||||||
|
type ctlServer struct {
|
||||||
|
l *slog.Logger
|
||||||
|
ctx context.Context
|
||||||
|
srv *diag.Server
|
||||||
|
|
||||||
|
runMu sync.Mutex
|
||||||
|
runCfg *ctlConfig
|
||||||
|
run *ctlRuntime
|
||||||
|
}
|
||||||
|
|
||||||
|
// ctlRuntime is the live state owned by a single Start invocation.
|
||||||
|
type ctlRuntime struct {
|
||||||
|
cancel context.CancelFunc
|
||||||
|
listener net.Listener
|
||||||
|
}
|
||||||
|
|
||||||
|
// newCtlServerFromConfig builds a ctlServer, parses the config, and registers a reload
|
||||||
|
// callback. It deliberately does not start listening: there is no interface yet, and
|
||||||
|
// Control.Start is what launches the first runtime. The callback is registered before the
|
||||||
|
// config is parsed so a SIGHUP can fix a bad block even if the first parse failed.
|
||||||
|
//
|
||||||
|
// reg is only held, never read, until Start runs. That is what lets this be constructed
|
||||||
|
// before attachCommands has populated the registry.
|
||||||
|
func newCtlServerFromConfig(ctx context.Context, l *slog.Logger, c *config.C, reg *diag.Registry) (*ctlServer, error) {
|
||||||
|
s := &ctlServer{
|
||||||
|
l: l,
|
||||||
|
ctx: ctx,
|
||||||
|
srv: diag.NewServer(l, reg),
|
||||||
|
}
|
||||||
|
|
||||||
|
c.RegisterReloadCallback(func(c *config.C) {
|
||||||
|
if err := s.reload(c, false); err != nil {
|
||||||
|
s.l.Error("Failed to reload ctl from config", "error", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
if err := s.reload(c, true); err != nil {
|
||||||
|
return s, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return s, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// loadCtlConfig parses and validates the `ctl` block. An empty socket path while enabled is
|
||||||
|
// not an error: it means the platform has no default and the operator did not name one, so
|
||||||
|
// there is simply nothing to listen on.
|
||||||
|
func loadCtlConfig(c *config.C) (ctlConfig, error) {
|
||||||
|
cfg := ctlConfig{
|
||||||
|
enabled: c.GetBool("ctl.enabled", true),
|
||||||
|
socket: c.GetString("ctl.socket", diag.DefaultSocketPath()),
|
||||||
|
explicit: c.IsSet("ctl.socket"),
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.enabled && cfg.socket != "" && !filepath.IsAbs(cfg.socket) {
|
||||||
|
return cfg, util.NewContextualError("ctl.socket must be an absolute path", m{"path": cfg.socket}, nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
return cfg, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// reload parses the config and records it, then reconciles the running listener against it:
|
||||||
|
//
|
||||||
|
// - newly enabled -> spawn Start
|
||||||
|
// - newly disabled -> Stop the runtime
|
||||||
|
// - socket moved (still enabled) -> Stop the old, Start the new
|
||||||
|
// - no change -> no-op
|
||||||
|
//
|
||||||
|
// On the initial call it only records configuration; Control.Start is what launches the first
|
||||||
|
// runtime via ctlStart. There is no interface to serve yet at that point.
|
||||||
|
func (s *ctlServer) reload(c *config.C, initial bool) error {
|
||||||
|
newCfg, err := loadCtlConfig(c)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
s.runMu.Lock()
|
||||||
|
sameCfg := s.runCfg != nil && *s.runCfg == newCfg
|
||||||
|
s.runCfg = &newCfg
|
||||||
|
running := s.run != nil
|
||||||
|
s.runMu.Unlock()
|
||||||
|
|
||||||
|
if initial || sameCfg {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if running {
|
||||||
|
s.Stop()
|
||||||
|
}
|
||||||
|
|
||||||
|
if newCfg.enabled && newCfg.socket != "" {
|
||||||
|
go s.Start()
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Start binds the socket and serves until Stop is called or ctx fires. Safe to call when ctl
|
||||||
|
// is disabled or already running: both no-op.
|
||||||
|
func (s *ctlServer) Start() {
|
||||||
|
s.runMu.Lock()
|
||||||
|
if s.ctx.Err() != nil || s.run != nil || s.runCfg == nil {
|
||||||
|
s.runMu.Unlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
cfg := *s.runCfg
|
||||||
|
s.runMu.Unlock()
|
||||||
|
|
||||||
|
if !cfg.enabled || cfg.socket == "" {
|
||||||
|
if cfg.enabled {
|
||||||
|
s.l.Info("ctl has no socket path on this platform, `nebula ctl` will not be available",
|
||||||
|
"hint", "set ctl.socket to enable it",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
listener, err := diag.Listen(cfg.socket)
|
||||||
|
if err != nil {
|
||||||
|
// A default path nebula cannot create is an ordinary state for an unprivileged
|
||||||
|
// install; a path the operator chose failing to bind is something they want to know
|
||||||
|
// about. Either way ctl is optional and nebula carries on without it.
|
||||||
|
if cfg.explicit {
|
||||||
|
s.l.Error("Failed to listen on the ctl socket", "ctlSocket", cfg.socket, "error", err)
|
||||||
|
} else {
|
||||||
|
s.l.Info("Not serving the ctl socket, `nebula ctl` will not be available",
|
||||||
|
"ctlSocket", cfg.socket,
|
||||||
|
"error", err,
|
||||||
|
"hint", "set ctl.socket to a path nebula can write, or ctl.enabled to false",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Drop the cached config so a SIGHUP retries once the underlying problem is fixed,
|
||||||
|
// even when the config itself is unchanged.
|
||||||
|
s.runMu.Lock()
|
||||||
|
if s.runCfg != nil && *s.runCfg == cfg {
|
||||||
|
s.runCfg = nil
|
||||||
|
}
|
||||||
|
s.runMu.Unlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
runCtx, cancel := context.WithCancel(s.ctx)
|
||||||
|
rt := &ctlRuntime{cancel: cancel, listener: listener}
|
||||||
|
|
||||||
|
s.runMu.Lock()
|
||||||
|
// Losing the race against a Stop or a competing Start means this listener is already
|
||||||
|
// obsolete. Close it rather than serving a socket nobody will tear down.
|
||||||
|
if s.ctx.Err() != nil || s.run != nil {
|
||||||
|
s.runMu.Unlock()
|
||||||
|
cancel()
|
||||||
|
_ = listener.Close()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.run = rt
|
||||||
|
s.runMu.Unlock()
|
||||||
|
|
||||||
|
s.l.Info("ctl socket is listening", "ctlSocket", cfg.socket)
|
||||||
|
|
||||||
|
err = s.srv.Serve(runCtx, listener)
|
||||||
|
if err != nil {
|
||||||
|
s.l.Error("The ctl listener stopped", "ctlSocket", cfg.socket, "error", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clear our runtime only if nothing has replaced it.
|
||||||
|
s.runMu.Lock()
|
||||||
|
if s.run == rt {
|
||||||
|
rt.cancel()
|
||||||
|
s.run = nil
|
||||||
|
if err != nil {
|
||||||
|
// An unclean exit leaves runCfg cached as if it were applied, so drop it and let a
|
||||||
|
// SIGHUP retry.
|
||||||
|
s.runCfg = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
s.runMu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Stop closes the listener and unlinks the socket. It deliberately does not touch connections
|
||||||
|
// that are already being served: `nebula ctl reload` runs every reload callback inline on its
|
||||||
|
// own connection, including this one, and hanging up on it would truncate the response to a
|
||||||
|
// reload that actually succeeded.
|
||||||
|
//
|
||||||
|
// The socket file is removed by net.UnixListener's unlink-on-close, so there is no os.Remove
|
||||||
|
// here; doing it by hand would delete a successor's socket after a fast reload.
|
||||||
|
func (s *ctlServer) Stop() {
|
||||||
|
s.runMu.Lock()
|
||||||
|
rt := s.run
|
||||||
|
s.run = nil
|
||||||
|
s.runMu.Unlock()
|
||||||
|
|
||||||
|
if rt == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
rt.cancel()
|
||||||
|
if err := rt.listener.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
|
||||||
|
s.l.Warn("Failed to close the ctl listener", "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
+305
@@ -0,0 +1,305 @@
|
|||||||
|
//go:build !windows
|
||||||
|
|
||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"log/slog"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/diag"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newTestCtlServer(t *testing.T) (*ctlServer, *config.C) {
|
||||||
|
t.Helper()
|
||||||
|
l := slog.New(slog.DiscardHandler)
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
t.Cleanup(cancel)
|
||||||
|
|
||||||
|
return &ctlServer{
|
||||||
|
l: l,
|
||||||
|
ctx: ctx,
|
||||||
|
srv: diag.NewServer(l, diag.NewRegistry()),
|
||||||
|
}, config.NewC(l)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setCtlConfig(c *config.C, m map[string]any) {
|
||||||
|
c.Settings["ctl"] = m
|
||||||
|
}
|
||||||
|
|
||||||
|
func currentCtlRuntime(s *ctlServer) *ctlRuntime {
|
||||||
|
s.runMu.Lock()
|
||||||
|
defer s.runMu.Unlock()
|
||||||
|
return s.run
|
||||||
|
}
|
||||||
|
|
||||||
|
// testCtlSocket returns a short socket path, see the note in diag/server_test.go about
|
||||||
|
// sun_path on darwin.
|
||||||
|
func testCtlSocket(t *testing.T) string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
dir, err := os.MkdirTemp("/tmp", "nebctl")
|
||||||
|
require.NoError(t, err)
|
||||||
|
t.Cleanup(func() { _ = os.RemoveAll(dir) })
|
||||||
|
|
||||||
|
return filepath.Join(dir, "ctl.sock")
|
||||||
|
}
|
||||||
|
|
||||||
|
func startCtl(t *testing.T, s *ctlServer) chan struct{} {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
s.Start()
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
return done
|
||||||
|
}
|
||||||
|
|
||||||
|
func requireCtlStopped(t *testing.T, done chan struct{}) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("ctl Start did not return after Stop")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCtlServer_loadConfig(t *testing.T) {
|
||||||
|
t.Run("defaults to enabled at the platform path", func(t *testing.T) {
|
||||||
|
_, c := newTestCtlServer(t)
|
||||||
|
|
||||||
|
cfg, err := loadCtlConfig(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.True(t, cfg.enabled)
|
||||||
|
assert.Equal(t, diag.DefaultSocketPath(), cfg.socket)
|
||||||
|
assert.False(t, cfg.explicit)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("an operator chosen path is recorded as explicit", func(t *testing.T) {
|
||||||
|
_, c := newTestCtlServer(t)
|
||||||
|
setCtlConfig(c, map[string]any{"socket": "/run/somewhere/ctl.sock"})
|
||||||
|
|
||||||
|
cfg, err := loadCtlConfig(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, "/run/somewhere/ctl.sock", cfg.socket)
|
||||||
|
assert.True(t, cfg.explicit)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a relative path is rejected", func(t *testing.T) {
|
||||||
|
_, c := newTestCtlServer(t)
|
||||||
|
setCtlConfig(c, map[string]any{"socket": "ctl.sock"})
|
||||||
|
|
||||||
|
_, err := loadCtlConfig(c)
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "must be an absolute path")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a relative path is not rejected when ctl is off", func(t *testing.T) {
|
||||||
|
_, c := newTestCtlServer(t)
|
||||||
|
setCtlConfig(c, map[string]any{"enabled": false, "socket": "ctl.sock"})
|
||||||
|
|
||||||
|
_, err := loadCtlConfig(c)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCtlServer_reload(t *testing.T) {
|
||||||
|
t.Run("the initial reload records config without listening", func(t *testing.T) {
|
||||||
|
s, c := newTestCtlServer(t)
|
||||||
|
setCtlConfig(c, map[string]any{"socket": testCtlSocket(t)})
|
||||||
|
|
||||||
|
require.NoError(t, s.reload(c, true))
|
||||||
|
assert.Nil(t, currentCtlRuntime(s), "Control.Start is what starts listening")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("enabling on reload starts listening", func(t *testing.T) {
|
||||||
|
s, c := newTestCtlServer(t)
|
||||||
|
path := testCtlSocket(t)
|
||||||
|
setCtlConfig(c, map[string]any{"enabled": false, "socket": path})
|
||||||
|
require.NoError(t, s.reload(c, true))
|
||||||
|
|
||||||
|
setCtlConfig(c, map[string]any{"enabled": true, "socket": path})
|
||||||
|
require.NoError(t, s.reload(c, false))
|
||||||
|
|
||||||
|
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||||
|
assert.FileExists(t, path)
|
||||||
|
|
||||||
|
s.Stop()
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("disabling on reload stops listening and unlinks", func(t *testing.T) {
|
||||||
|
s, c := newTestCtlServer(t)
|
||||||
|
path := testCtlSocket(t)
|
||||||
|
setCtlConfig(c, map[string]any{"enabled": true, "socket": path})
|
||||||
|
require.NoError(t, s.reload(c, true))
|
||||||
|
|
||||||
|
done := startCtl(t, s)
|
||||||
|
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||||
|
|
||||||
|
setCtlConfig(c, map[string]any{"enabled": false, "socket": path})
|
||||||
|
require.NoError(t, s.reload(c, false))
|
||||||
|
|
||||||
|
requireCtlStopped(t, done)
|
||||||
|
assert.Nil(t, currentCtlRuntime(s))
|
||||||
|
assert.NoFileExists(t, path)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("moving the socket restarts at the new path", func(t *testing.T) {
|
||||||
|
s, c := newTestCtlServer(t)
|
||||||
|
oldPath := testCtlSocket(t)
|
||||||
|
newPath := testCtlSocket(t)
|
||||||
|
setCtlConfig(c, map[string]any{"socket": oldPath})
|
||||||
|
require.NoError(t, s.reload(c, true))
|
||||||
|
|
||||||
|
done := startCtl(t, s)
|
||||||
|
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||||
|
require.FileExists(t, oldPath)
|
||||||
|
|
||||||
|
setCtlConfig(c, map[string]any{"socket": newPath})
|
||||||
|
require.NoError(t, s.reload(c, false))
|
||||||
|
requireCtlStopped(t, done)
|
||||||
|
|
||||||
|
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||||
|
assert.FileExists(t, newPath)
|
||||||
|
assert.NoFileExists(t, oldPath, "the old socket should have been unlinked")
|
||||||
|
|
||||||
|
s.Stop()
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("an unchanged config leaves the listener alone", func(t *testing.T) {
|
||||||
|
s, c := newTestCtlServer(t)
|
||||||
|
setCtlConfig(c, map[string]any{"socket": testCtlSocket(t)})
|
||||||
|
require.NoError(t, s.reload(c, true))
|
||||||
|
|
||||||
|
startCtl(t, s)
|
||||||
|
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||||
|
before := currentCtlRuntime(s)
|
||||||
|
|
||||||
|
require.NoError(t, s.reload(c, false))
|
||||||
|
assert.Same(t, before, currentCtlRuntime(s), "the runtime should not have been replaced")
|
||||||
|
|
||||||
|
s.Stop()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCtlServer_Start(t *testing.T) {
|
||||||
|
t.Run("a command can be run over the socket", func(t *testing.T) {
|
||||||
|
l := slog.New(slog.DiscardHandler)
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
t.Cleanup(cancel)
|
||||||
|
|
||||||
|
reg := diag.NewRegistry()
|
||||||
|
s := &ctlServer{l: l, ctx: ctx, srv: diag.NewServer(l, reg)}
|
||||||
|
c := config.NewC(l)
|
||||||
|
|
||||||
|
path := testCtlSocket(t)
|
||||||
|
setCtlConfig(c, map[string]any{"socket": path})
|
||||||
|
require.NoError(t, s.reload(c, true))
|
||||||
|
|
||||||
|
startCtl(t, s)
|
||||||
|
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||||
|
|
||||||
|
client, err := diag.Dial(path)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer client.Close()
|
||||||
|
|
||||||
|
out := &testWriter{}
|
||||||
|
status, err := client.Run([]string{"help"}, out)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, diag.StatusOK, status)
|
||||||
|
assert.Contains(t, out.String(), "Available commands:")
|
||||||
|
|
||||||
|
s.Stop()
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Start is a no-op when ctl is disabled", func(t *testing.T) {
|
||||||
|
s, c := newTestCtlServer(t)
|
||||||
|
setCtlConfig(c, map[string]any{"enabled": false, "socket": testCtlSocket(t)})
|
||||||
|
require.NoError(t, s.reload(c, true))
|
||||||
|
|
||||||
|
s.Start()
|
||||||
|
assert.Nil(t, currentCtlRuntime(s))
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Start is a no-op with no socket path for this platform", func(t *testing.T) {
|
||||||
|
s, c := newTestCtlServer(t)
|
||||||
|
setCtlConfig(c, map[string]any{"enabled": true, "socket": ""})
|
||||||
|
require.NoError(t, s.reload(c, true))
|
||||||
|
|
||||||
|
s.Start()
|
||||||
|
assert.Nil(t, currentCtlRuntime(s))
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Start is a no-op after the context is cancelled", func(t *testing.T) {
|
||||||
|
l := slog.New(slog.DiscardHandler)
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
s := &ctlServer{l: l, ctx: ctx, srv: diag.NewServer(l, diag.NewRegistry())}
|
||||||
|
c := config.NewC(l)
|
||||||
|
|
||||||
|
path := testCtlSocket(t)
|
||||||
|
setCtlConfig(c, map[string]any{"socket": path})
|
||||||
|
require.NoError(t, s.reload(c, true))
|
||||||
|
cancel()
|
||||||
|
|
||||||
|
s.Start()
|
||||||
|
assert.Nil(t, currentCtlRuntime(s))
|
||||||
|
assert.NoFileExists(t, path)
|
||||||
|
})
|
||||||
|
|
||||||
|
// A path nebula cannot bind must not stop it from running, and a SIGHUP with the same
|
||||||
|
// config has to be able to retry once the problem is fixed.
|
||||||
|
t.Run("a listen failure is survivable and retried on the next reload", func(t *testing.T) {
|
||||||
|
s, c := newTestCtlServer(t)
|
||||||
|
path := testCtlSocket(t)
|
||||||
|
require.NoError(t, os.MkdirAll(filepath.Dir(path), 0700))
|
||||||
|
require.NoError(t, os.WriteFile(path, []byte("in the way"), 0600))
|
||||||
|
|
||||||
|
setCtlConfig(c, map[string]any{"socket": path})
|
||||||
|
require.NoError(t, s.reload(c, true))
|
||||||
|
|
||||||
|
s.Start()
|
||||||
|
assert.Nil(t, currentCtlRuntime(s))
|
||||||
|
|
||||||
|
s.runMu.Lock()
|
||||||
|
cachedCfg := s.runCfg
|
||||||
|
s.runMu.Unlock()
|
||||||
|
assert.Nil(t, cachedCfg, "the cached config should be dropped so a reload retries")
|
||||||
|
|
||||||
|
require.NoError(t, os.Remove(path))
|
||||||
|
require.NoError(t, s.reload(c, false))
|
||||||
|
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||||
|
|
||||||
|
s.Stop()
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Stop is idempotent", func(t *testing.T) {
|
||||||
|
s, c := newTestCtlServer(t)
|
||||||
|
setCtlConfig(c, map[string]any{"socket": testCtlSocket(t)})
|
||||||
|
require.NoError(t, s.reload(c, true))
|
||||||
|
|
||||||
|
done := startCtl(t, s)
|
||||||
|
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||||
|
|
||||||
|
s.Stop()
|
||||||
|
requireCtlStopped(t, done)
|
||||||
|
assert.NotPanics(t, s.Stop)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// testWriter collects command output.
|
||||||
|
type testWriter struct{ b []byte }
|
||||||
|
|
||||||
|
func (w *testWriter) Write(p []byte) (int, error) {
|
||||||
|
w.b = append(w.b, p...)
|
||||||
|
return len(p), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *testWriter) String() string { return string(w.b) }
|
||||||
@@ -0,0 +1,41 @@
|
|||||||
|
package diag
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"io"
|
||||||
|
"net"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// dialTimeout bounds the connect only. A command may take as long as it likes to answer.
|
||||||
|
const dialTimeout = 2 * time.Second
|
||||||
|
|
||||||
|
// Client is a connection to a nebula serving the ctl socket. It carries exactly one command.
|
||||||
|
type Client struct {
|
||||||
|
conn net.Conn
|
||||||
|
}
|
||||||
|
|
||||||
|
// Dial connects to the nebula serving at path. On a platform without socket support the
|
||||||
|
// returned error wraps ErrNotSupported.
|
||||||
|
func Dial(path string) (*Client, error) {
|
||||||
|
conn, err := dialSocket(path, dialTimeout)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return &Client{conn: conn}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Run sends args and streams the command's output to out, returning the command's exit
|
||||||
|
// status. A non-nil error means the exchange itself failed and the status means nothing.
|
||||||
|
func (c *Client) Run(args []string, out io.Writer) (int, error) {
|
||||||
|
if err := writeRequest(c.conn, args); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return readResponse(bufio.NewReader(c.conn), out)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) Close() error {
|
||||||
|
return c.conn.Close()
|
||||||
|
}
|
||||||
@@ -1,4 +1,4 @@
|
|||||||
package sshd
|
package diag
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
@@ -10,6 +10,17 @@ import (
|
|||||||
"github.com/armon/go-radix"
|
"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
|
// 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
|
// It should return a flag.FlagSet instance and a pointer to the struct that will contain parsed flags
|
||||||
type CommandFlags func() (*flag.FlagSet, any)
|
type CommandFlags func() (*flag.FlagSet, any)
|
||||||
@@ -44,8 +55,10 @@ func execCommand(c *Command, args []string, w StringWriter) error {
|
|||||||
fl.SetOutput(w.GetWriter())
|
fl.SetOutput(w.GetWriter())
|
||||||
err := fl.Parse(args)
|
err := fl.Parse(args)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// fl.Parse has dumped error information to the user via the w writer.
|
// fl.Parse has dumped error information to the user via the w writer, so
|
||||||
return err
|
// the wrapper exists purely so a transport can tell a usage problem from a
|
||||||
|
// command that ran and failed.
|
||||||
|
return fmt.Errorf("%w: %w", ErrUsage, err)
|
||||||
}
|
}
|
||||||
args = fl.Args()
|
args = fl.Args()
|
||||||
}
|
}
|
||||||
+226
@@ -0,0 +1,226 @@
|
|||||||
|
package diag
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"encoding/binary"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The ctl protocol is one request, one response, one connection.
|
||||||
|
//
|
||||||
|
// The request is a single JSON line. argv travels as a list rather than a joined string so
|
||||||
|
// that a path with a space in it survives the trip; the client already has a real argv from
|
||||||
|
// the operating system and re-splitting it would only ever lose information.
|
||||||
|
//
|
||||||
|
// The response is a stream of frames rather than raw bytes followed by a status line,
|
||||||
|
// because there is no sentinel that is safe to look for: `print-cert -raw` emits arbitrary
|
||||||
|
// PEM and `list-hostmap -json` emits arbitrary JSON, either of which could contain whatever
|
||||||
|
// terminator we picked.
|
||||||
|
const (
|
||||||
|
// ProtoVersion is the only request version this build understands. An unknown version
|
||||||
|
// gets a legible error rather than a hang, which is the whole point of sending it.
|
||||||
|
ProtoVersion = 1
|
||||||
|
|
||||||
|
// frameOutput carries raw command output, destined for the client's stdout.
|
||||||
|
frameOutput = 0x01
|
||||||
|
// frameEnd carries a JSON endPayload and is the last frame on a connection.
|
||||||
|
frameEnd = 0x02
|
||||||
|
// frameStderr is reserved. Commands write to a single writer today, so there is nothing
|
||||||
|
// to put in it, but holding the number means adding one later needs no version bump.
|
||||||
|
frameStderr = 0x03
|
||||||
|
|
||||||
|
// maxFrame bounds a single frame's payload. Larger writes are split across frames.
|
||||||
|
maxFrame = 64 * 1024
|
||||||
|
// maxRequest bounds the request line, so a client that never sends a newline cannot make
|
||||||
|
// nebula buffer without limit.
|
||||||
|
maxRequest = 64 * 1024
|
||||||
|
// outputBuffer is what keeps json.NewEncoder(w.GetWriter()) from emitting a frame per
|
||||||
|
// token; output accumulates here and flushes in useful sized chunks.
|
||||||
|
outputBuffer = 32 * 1024
|
||||||
|
)
|
||||||
|
|
||||||
|
// ErrTruncated means the connection ended before the end frame arrived, which is how a
|
||||||
|
// client notices that nebula died or was torn down partway through a command.
|
||||||
|
var ErrTruncated = errors.New("connection closed before the command finished")
|
||||||
|
|
||||||
|
// request is the JSON line a client sends.
|
||||||
|
type request struct {
|
||||||
|
Version int `json:"version"`
|
||||||
|
Args []string `json:"args"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// endPayload is the JSON body of the end frame. Error is set only when Status is non-zero
|
||||||
|
// and describes a failure to run the command, not a failure the command itself reported.
|
||||||
|
type endPayload struct {
|
||||||
|
Status int `json:"status"`
|
||||||
|
Error string `json:"error,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeRequest sends the request line.
|
||||||
|
func writeRequest(w io.Writer, args []string) error {
|
||||||
|
b, err := json.Marshal(request{Version: ProtoVersion, Args: args})
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(b)+1 > maxRequest {
|
||||||
|
return fmt.Errorf("command line is too long: %d bytes", len(b))
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = w.Write(append(b, '\n'))
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// readRequest reads and validates one request line.
|
||||||
|
func readRequest(r *bufio.Reader) (request, error) {
|
||||||
|
var req request
|
||||||
|
|
||||||
|
line, err := readLimitedLine(r, maxRequest)
|
||||||
|
if err != nil {
|
||||||
|
return req, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal(line, &req); err != nil {
|
||||||
|
return req, fmt.Errorf("malformed request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if req.Version != ProtoVersion {
|
||||||
|
return req, fmt.Errorf("unsupported protocol version %d, this nebula speaks version %d", req.Version, ProtoVersion)
|
||||||
|
}
|
||||||
|
|
||||||
|
return req, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// readLimitedLine reads through the next newline, refusing a line longer than limit rather
|
||||||
|
// than buffering whatever an unfriendly client decides to send.
|
||||||
|
func readLimitedLine(r *bufio.Reader, limit int) ([]byte, error) {
|
||||||
|
line := make([]byte, 0, 256)
|
||||||
|
for {
|
||||||
|
b, err := r.ReadByte()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if b == '\n' {
|
||||||
|
return line, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(line) >= limit {
|
||||||
|
return nil, fmt.Errorf("request exceeded %d bytes without a newline", limit)
|
||||||
|
}
|
||||||
|
|
||||||
|
line = append(line, b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// frameWriter turns writes into output frames. It is handed to commands wrapped in a
|
||||||
|
// bufio.Writer, so a command that makes many small writes does not make many small frames.
|
||||||
|
type frameWriter struct {
|
||||||
|
w io.Writer
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *frameWriter) Write(b []byte) (int, error) {
|
||||||
|
written := 0
|
||||||
|
for {
|
||||||
|
chunk := b[written:]
|
||||||
|
if len(chunk) > maxFrame {
|
||||||
|
chunk = chunk[:maxFrame]
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := writeFrame(f.w, frameOutput, chunk); err != nil {
|
||||||
|
return written, err
|
||||||
|
}
|
||||||
|
|
||||||
|
written += len(chunk)
|
||||||
|
if written == len(b) {
|
||||||
|
return written, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeFrame emits one frame: a type byte, a big endian length, then the payload.
|
||||||
|
func writeFrame(w io.Writer, kind byte, payload []byte) error {
|
||||||
|
var hdr [5]byte
|
||||||
|
hdr[0] = kind
|
||||||
|
binary.BigEndian.PutUint32(hdr[1:], uint32(len(payload)))
|
||||||
|
|
||||||
|
if _, err := w.Write(hdr[:]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(payload) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := w.Write(payload)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeEnd emits the final frame. A transport error here is unreportable by definition, the
|
||||||
|
// connection is the only channel we have.
|
||||||
|
func writeEnd(w io.Writer, status int, msg string) error {
|
||||||
|
b, err := json.Marshal(endPayload{Status: status, Error: msg})
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return writeFrame(w, frameEnd, b)
|
||||||
|
}
|
||||||
|
|
||||||
|
// readResponse consumes frames until the end frame, copying output to out. It returns the
|
||||||
|
// command's exit status. A non-nil error means the exchange failed and the status is
|
||||||
|
// meaningless.
|
||||||
|
func readResponse(r io.Reader, out io.Writer) (int, error) {
|
||||||
|
var hdr [5]byte
|
||||||
|
|
||||||
|
for {
|
||||||
|
if _, err := io.ReadFull(r, hdr[:]); err != nil {
|
||||||
|
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) {
|
||||||
|
return 0, ErrTruncated
|
||||||
|
}
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
length := binary.BigEndian.Uint32(hdr[1:])
|
||||||
|
if length > maxFrame {
|
||||||
|
return 0, fmt.Errorf("frame of %d bytes exceeds the %d byte maximum", length, maxFrame)
|
||||||
|
}
|
||||||
|
|
||||||
|
payload := make([]byte, length)
|
||||||
|
if _, err := io.ReadFull(r, payload); err != nil {
|
||||||
|
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) {
|
||||||
|
return 0, ErrTruncated
|
||||||
|
}
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
switch hdr[0] {
|
||||||
|
case frameOutput:
|
||||||
|
if _, err := out.Write(payload); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
case frameEnd:
|
||||||
|
var end endPayload
|
||||||
|
if err := json.Unmarshal(payload, &end); err != nil {
|
||||||
|
return 0, fmt.Errorf("malformed end frame: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if end.Error != "" {
|
||||||
|
return end.Status, errors.New(end.Error)
|
||||||
|
}
|
||||||
|
|
||||||
|
return end.Status, nil
|
||||||
|
|
||||||
|
case frameStderr:
|
||||||
|
// Reserved and unused by this build. Skipping rather than failing means an older
|
||||||
|
// client stays usable against a newer nebula that starts sending them.
|
||||||
|
|
||||||
|
default:
|
||||||
|
return 0, fmt.Errorf("unknown frame type 0x%02x", hdr[0])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,144 @@
|
|||||||
|
package diag
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"bytes"
|
||||||
|
"encoding/binary"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestRequestRoundTrip(t *testing.T) {
|
||||||
|
t.Run("argv survives a round trip, spaces and all", func(t *testing.T) {
|
||||||
|
buf := &bytes.Buffer{}
|
||||||
|
args := []string{"start-cpu-profile", "/tmp/a path.pb.gz", "-json"}
|
||||||
|
require.NoError(t, writeRequest(buf, args))
|
||||||
|
|
||||||
|
req, err := readRequest(bufio.NewReader(buf))
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, ProtoVersion, req.Version)
|
||||||
|
assert.Equal(t, args, req.Args)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("an unknown version is refused by name", func(t *testing.T) {
|
||||||
|
r := bufio.NewReader(strings.NewReader(`{"version":99,"args":["version"]}` + "\n"))
|
||||||
|
|
||||||
|
_, err := readRequest(r)
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "unsupported protocol version 99")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("malformed json is refused", func(t *testing.T) {
|
||||||
|
r := bufio.NewReader(strings.NewReader("not json\n"))
|
||||||
|
|
||||||
|
_, err := readRequest(r)
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "malformed request")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a line without a newline is bounded rather than buffered forever", func(t *testing.T) {
|
||||||
|
r := bufio.NewReader(strings.NewReader(strings.Repeat("a", maxRequest+10)))
|
||||||
|
|
||||||
|
_, err := readRequest(r)
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "without a newline")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResponseRoundTrip(t *testing.T) {
|
||||||
|
t.Run("output and status survive a round trip", func(t *testing.T) {
|
||||||
|
wire := &bytes.Buffer{}
|
||||||
|
w := bufio.NewWriterSize(&frameWriter{w: wire}, outputBuffer)
|
||||||
|
require.NoError(t, NewWriter(w).WriteLine("hello"))
|
||||||
|
require.NoError(t, w.Flush())
|
||||||
|
require.NoError(t, writeEnd(wire, StatusOK, ""))
|
||||||
|
|
||||||
|
out := &bytes.Buffer{}
|
||||||
|
status, err := readResponse(wire, out)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, StatusOK, status)
|
||||||
|
assert.Equal(t, "hello\n", out.String())
|
||||||
|
})
|
||||||
|
|
||||||
|
// print-cert -raw and list-hostmap -json both emit arbitrary bytes, so a payload larger
|
||||||
|
// than one frame has to reassemble exactly.
|
||||||
|
t.Run("a payload larger than one frame reassembles byte for byte", func(t *testing.T) {
|
||||||
|
big := bytes.Repeat([]byte("nebula"), maxFrame)
|
||||||
|
|
||||||
|
wire := &bytes.Buffer{}
|
||||||
|
fw := &frameWriter{w: wire}
|
||||||
|
n, err := fw.Write(big)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, len(big), n)
|
||||||
|
require.NoError(t, writeEnd(wire, StatusOK, ""))
|
||||||
|
|
||||||
|
out := &bytes.Buffer{}
|
||||||
|
status, err := readResponse(wire, out)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, StatusOK, status)
|
||||||
|
assert.Equal(t, big, out.Bytes())
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a non-zero status carries its message", func(t *testing.T) {
|
||||||
|
wire := &bytes.Buffer{}
|
||||||
|
require.NoError(t, writeEnd(wire, StatusError, "it went wrong"))
|
||||||
|
|
||||||
|
status, err := readResponse(wire, &bytes.Buffer{})
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Equal(t, StatusError, status)
|
||||||
|
assert.Contains(t, err.Error(), "it went wrong")
|
||||||
|
})
|
||||||
|
|
||||||
|
// This is how the CLI notices a nebula that died mid-command rather than silently
|
||||||
|
// reporting whatever partial output it managed to read.
|
||||||
|
t.Run("a stream ending without an end frame is truncated, not successful", func(t *testing.T) {
|
||||||
|
wire := &bytes.Buffer{}
|
||||||
|
_, err := (&frameWriter{w: wire}).Write([]byte("partial"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
out := &bytes.Buffer{}
|
||||||
|
_, err = readResponse(wire, out)
|
||||||
|
assert.ErrorIs(t, err, ErrTruncated)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a truncated frame header is truncated, not successful", func(t *testing.T) {
|
||||||
|
_, err := readResponse(bytes.NewReader([]byte{frameOutput, 0x00}), &bytes.Buffer{})
|
||||||
|
assert.ErrorIs(t, err, ErrTruncated)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("an oversized frame is refused rather than allocated", func(t *testing.T) {
|
||||||
|
var hdr [5]byte
|
||||||
|
hdr[0] = frameOutput
|
||||||
|
binary.BigEndian.PutUint32(hdr[1:], maxFrame+1)
|
||||||
|
|
||||||
|
_, err := readResponse(bytes.NewReader(hdr[:]), &bytes.Buffer{})
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "exceeds")
|
||||||
|
})
|
||||||
|
|
||||||
|
// A reserved frame an older client does not understand must not break it.
|
||||||
|
t.Run("a reserved frame type is skipped", func(t *testing.T) {
|
||||||
|
wire := &bytes.Buffer{}
|
||||||
|
require.NoError(t, writeFrame(wire, frameStderr, []byte("future")))
|
||||||
|
require.NoError(t, writeFrame(wire, frameOutput, []byte("now")))
|
||||||
|
require.NoError(t, writeEnd(wire, StatusOK, ""))
|
||||||
|
|
||||||
|
out := &bytes.Buffer{}
|
||||||
|
status, err := readResponse(wire, out)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, StatusOK, status)
|
||||||
|
assert.Equal(t, "now", out.String())
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("an unknown frame type is an error", func(t *testing.T) {
|
||||||
|
wire := &bytes.Buffer{}
|
||||||
|
require.NoError(t, writeFrame(wire, 0x7f, nil))
|
||||||
|
|
||||||
|
_, err := readResponse(wire, &bytes.Buffer{})
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "unknown frame type")
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -0,0 +1,125 @@
|
|||||||
|
package diag
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/anmitsu/go-shlex"
|
||||||
|
"github.com/armon/go-radix"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Registry is the set of commands nebula exposes for debugging and administration. It is
|
||||||
|
// transport neutral: the ssh console and the `nebula ctl` unix socket dispatch against the
|
||||||
|
// same registry, and neither knows the other exists.
|
||||||
|
//
|
||||||
|
// Registration is expected to happen once during startup, before any transport is serving,
|
||||||
|
// but the lock makes a late RegisterCommand safe rather than a data race waiting to happen.
|
||||||
|
type Registry struct {
|
||||||
|
mu sync.RWMutex
|
||||||
|
commands *radix.Tree
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewRegistry returns a registry containing only `help`. Everything else is attached by
|
||||||
|
// the caller, see attachCommands in the nebula package.
|
||||||
|
func NewRegistry() *Registry {
|
||||||
|
r := &Registry{commands: radix.New()}
|
||||||
|
|
||||||
|
r.RegisterCommand(&Command{
|
||||||
|
Name: "help",
|
||||||
|
ShortDescription: "prints available commands or help <command> for specific usage info",
|
||||||
|
Callback: func(a any, args []string, w StringWriter) error {
|
||||||
|
return r.help(args, w)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterCommand adds a command that a user can run.
|
||||||
|
func (r *Registry) RegisterCommand(c *Command) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
r.commands.Insert(c.Name, c)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clone returns an independent copy sharing no tree with the original. The ssh session uses
|
||||||
|
// this so the `logout` command it adds for itself is invisible to every other session, and
|
||||||
|
// to `nebula ctl`.
|
||||||
|
func (r *Registry) Clone() *Registry {
|
||||||
|
r.mu.RLock()
|
||||||
|
defer r.mu.RUnlock()
|
||||||
|
return &Registry{commands: radix.NewFromMap(r.commands.ToMap())}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Match returns every registered command name carrying the given prefix, for tab completion.
|
||||||
|
func (r *Registry) Match(prefix string) []string {
|
||||||
|
r.mu.RLock()
|
||||||
|
defer r.mu.RUnlock()
|
||||||
|
return matchCommand(r.commands, prefix)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Dispatch splits line the way a shell would and runs the result. The ssh console uses this
|
||||||
|
// because a terminal only ever hands it a line; a transport that already has a real argv
|
||||||
|
// should call DispatchArgs instead rather than round tripping through a quoting parser.
|
||||||
|
func (r *Registry) Dispatch(line string, w StringWriter) error {
|
||||||
|
args, err := shlex.Split(line, true)
|
||||||
|
if err != nil {
|
||||||
|
if wErr := w.WriteLine(fmt.Sprintf("Unable to parse command: %s", err)); wErr != nil {
|
||||||
|
return wErr
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return r.DispatchArgs(args, w)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DispatchArgs runs args[0] with args[1:] as its arguments, writing everything the command
|
||||||
|
// produces to w. An empty args dumps the command list, matching what an empty line does on
|
||||||
|
// the ssh console.
|
||||||
|
//
|
||||||
|
// Callbacks report user facing problems as prose on w and return nil by convention, so a
|
||||||
|
// non-nil error here means the command could not be run at all: ErrUnknownCommand, an
|
||||||
|
// ErrUsage wrapped flag failure, or an internal failure a callback chose to surface.
|
||||||
|
func (r *Registry) DispatchArgs(args []string, w StringWriter) error {
|
||||||
|
if len(args) == 0 {
|
||||||
|
r.mu.RLock()
|
||||||
|
defer r.mu.RUnlock()
|
||||||
|
dumpCommands(r.commands, w)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
r.mu.RLock()
|
||||||
|
cmd, err := lookupCommand(r.commands, args[0])
|
||||||
|
r.mu.RUnlock()
|
||||||
|
if err != nil {
|
||||||
|
if wErr := w.WriteLine(fmt.Sprintf("Command lookup failed: %s", err)); wErr != nil {
|
||||||
|
return wErr
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if cmd == nil {
|
||||||
|
if wErr := w.WriteLine(fmt.Sprintf("Did not understand: %s", args[0])); wErr != nil {
|
||||||
|
return wErr
|
||||||
|
}
|
||||||
|
r.mu.RLock()
|
||||||
|
defer r.mu.RUnlock()
|
||||||
|
dumpCommands(r.commands, w)
|
||||||
|
return fmt.Errorf("%w: %s", ErrUnknownCommand, args[0])
|
||||||
|
}
|
||||||
|
|
||||||
|
// -h and -help anywhere in the arguments mean the user wants to know how the command
|
||||||
|
// works, not to run it.
|
||||||
|
if checkHelpArgs(args) {
|
||||||
|
return r.help([]string{cmd.Name}, w)
|
||||||
|
}
|
||||||
|
|
||||||
|
return execCommand(cmd, args[1:], w)
|
||||||
|
}
|
||||||
|
|
||||||
|
// help renders the command list, or one command's usage, onto w.
|
||||||
|
func (r *Registry) help(args []string, w StringWriter) error {
|
||||||
|
r.mu.RLock()
|
||||||
|
defer r.mu.RUnlock()
|
||||||
|
return helpCallback(r.commands, args, w)
|
||||||
|
}
|
||||||
@@ -0,0 +1,167 @@
|
|||||||
|
package diag
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"flag"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
type testFlags struct {
|
||||||
|
Json bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// testCommand builds a command carrying a flag set, recording what the callback was actually
|
||||||
|
// handed so a test can assert on it.
|
||||||
|
func testCommand(name string, seen *any, args *[]string) *Command {
|
||||||
|
return &Command{
|
||||||
|
Name: name,
|
||||||
|
ShortDescription: name + " short description",
|
||||||
|
Flags: func() (*flag.FlagSet, any) {
|
||||||
|
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||||
|
f := &testFlags{}
|
||||||
|
fl.BoolVar(&f.Json, "json", false, "outputs json")
|
||||||
|
return fl, f
|
||||||
|
},
|
||||||
|
Callback: func(fs any, a []string, w StringWriter) error {
|
||||||
|
if seen != nil {
|
||||||
|
*seen = fs
|
||||||
|
}
|
||||||
|
if args != nil {
|
||||||
|
*args = a
|
||||||
|
}
|
||||||
|
return w.WriteLine("ran " + name)
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTestRegistry(t *testing.T) (*Registry, *bytes.Buffer, StringWriter) {
|
||||||
|
t.Helper()
|
||||||
|
buf := &bytes.Buffer{}
|
||||||
|
return NewRegistry(), buf, NewWriter(buf)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegistryDispatch(t *testing.T) {
|
||||||
|
t.Run("a new registry knows help and nothing else", func(t *testing.T) {
|
||||||
|
r, buf, w := newTestRegistry(t)
|
||||||
|
|
||||||
|
require.NoError(t, r.DispatchArgs([]string{"help"}, w))
|
||||||
|
assert.Contains(t, buf.String(), "help -")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("empty args dump the command list, matching an empty line on the console", func(t *testing.T) {
|
||||||
|
r, buf, w := newTestRegistry(t)
|
||||||
|
r.RegisterCommand(testCommand("do-thing", nil, nil))
|
||||||
|
|
||||||
|
require.NoError(t, r.DispatchArgs(nil, w))
|
||||||
|
assert.Contains(t, buf.String(), "Available commands:")
|
||||||
|
assert.Contains(t, buf.String(), "do-thing - do-thing short description")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("an unknown command reports ErrUnknownCommand and still tells the user", func(t *testing.T) {
|
||||||
|
r, buf, w := newTestRegistry(t)
|
||||||
|
|
||||||
|
err := r.DispatchArgs([]string{"nope"}, w)
|
||||||
|
require.ErrorIs(t, err, ErrUnknownCommand)
|
||||||
|
assert.Contains(t, buf.String(), "Did not understand: nope")
|
||||||
|
assert.Contains(t, buf.String(), "Available commands:")
|
||||||
|
})
|
||||||
|
|
||||||
|
// This is the hazard the ctl transport has to preserve: every callback in ssh.go begins by
|
||||||
|
// type asserting fs to its own concrete flags struct. Reach a callback without going
|
||||||
|
// through Command.Flags and every one of them fails.
|
||||||
|
t.Run("a callback is handed the concrete struct its Flags callback returned", func(t *testing.T) {
|
||||||
|
var seen any
|
||||||
|
r, _, w := newTestRegistry(t)
|
||||||
|
r.RegisterCommand(testCommand("do-thing", &seen, nil))
|
||||||
|
|
||||||
|
require.NoError(t, r.DispatchArgs([]string{"do-thing", "-json"}, w))
|
||||||
|
|
||||||
|
flags, ok := seen.(*testFlags)
|
||||||
|
require.True(t, ok, "callback was handed %T, not *testFlags", seen)
|
||||||
|
assert.True(t, flags.Json)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("positional arguments survive flag parsing", func(t *testing.T) {
|
||||||
|
var args []string
|
||||||
|
r, _, w := newTestRegistry(t)
|
||||||
|
r.RegisterCommand(testCommand("do-thing", nil, &args))
|
||||||
|
|
||||||
|
require.NoError(t, r.DispatchArgs([]string{"do-thing", "-json", "10.0.0.1"}, w))
|
||||||
|
assert.Equal(t, []string{"10.0.0.1"}, args)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Documents stdlib flag behaviour rather than endorsing it: parsing stops at the first
|
||||||
|
// positional, so a flag written after one is silently a positional too.
|
||||||
|
t.Run("a flag after a positional is not parsed as a flag", func(t *testing.T) {
|
||||||
|
var seen any
|
||||||
|
var args []string
|
||||||
|
r, _, w := newTestRegistry(t)
|
||||||
|
r.RegisterCommand(testCommand("do-thing", &seen, &args))
|
||||||
|
|
||||||
|
require.NoError(t, r.DispatchArgs([]string{"do-thing", "10.0.0.1", "-json"}, w))
|
||||||
|
assert.False(t, seen.(*testFlags).Json)
|
||||||
|
assert.Equal(t, []string{"10.0.0.1", "-json"}, args)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a bad flag reports ErrUsage and writes the usage text", func(t *testing.T) {
|
||||||
|
r, buf, w := newTestRegistry(t)
|
||||||
|
r.RegisterCommand(testCommand("do-thing", nil, nil))
|
||||||
|
|
||||||
|
err := r.DispatchArgs([]string{"do-thing", "-nope"}, w)
|
||||||
|
require.ErrorIs(t, err, ErrUsage)
|
||||||
|
assert.Contains(t, buf.String(), "flag provided but not defined")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("-h anywhere routes to help instead of running the command", func(t *testing.T) {
|
||||||
|
var seen any
|
||||||
|
r, buf, w := newTestRegistry(t)
|
||||||
|
r.RegisterCommand(testCommand("do-thing", &seen, nil))
|
||||||
|
|
||||||
|
require.NoError(t, r.DispatchArgs([]string{"do-thing", "-h"}, w))
|
||||||
|
assert.Nil(t, seen, "the callback should not have run")
|
||||||
|
assert.Contains(t, buf.String(), "do-thing - do-thing short description")
|
||||||
|
assert.Contains(t, buf.String(), "-json")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Dispatch splits a line the way a shell would", func(t *testing.T) {
|
||||||
|
var args []string
|
||||||
|
r, _, w := newTestRegistry(t)
|
||||||
|
r.RegisterCommand(testCommand("do-thing", nil, &args))
|
||||||
|
|
||||||
|
require.NoError(t, r.Dispatch(`do-thing "/tmp/a path.pb.gz"`, w))
|
||||||
|
assert.Equal(t, []string{"/tmp/a path.pb.gz"}, args)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Match returns names by prefix for tab completion", func(t *testing.T) {
|
||||||
|
r, _, _ := newTestRegistry(t)
|
||||||
|
r.RegisterCommand(testCommand("print-cert", nil, nil))
|
||||||
|
r.RegisterCommand(testCommand("print-tunnel", nil, nil))
|
||||||
|
r.RegisterCommand(testCommand("version", nil, nil))
|
||||||
|
|
||||||
|
assert.Equal(t, []string{"print-cert", "print-tunnel"}, r.Match("print-"))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// A clone is what keeps the ssh session's `logout` command from being visible to every other
|
||||||
|
// session, and to nebula ctl.
|
||||||
|
func TestRegistryCloneIsolation(t *testing.T) {
|
||||||
|
parent, _, w := newTestRegistry(t)
|
||||||
|
parent.RegisterCommand(testCommand("shared", nil, nil))
|
||||||
|
|
||||||
|
child := parent.Clone()
|
||||||
|
child.RegisterCommand(testCommand("logout", nil, nil))
|
||||||
|
|
||||||
|
require.NoError(t, child.DispatchArgs([]string{"logout"}, w))
|
||||||
|
|
||||||
|
buf := &bytes.Buffer{}
|
||||||
|
err := parent.DispatchArgs([]string{"logout"}, NewWriter(buf))
|
||||||
|
assert.ErrorIs(t, err, ErrUnknownCommand)
|
||||||
|
|
||||||
|
buf.Reset()
|
||||||
|
require.NoError(t, child.DispatchArgs([]string{"shared"}, NewWriter(buf)))
|
||||||
|
assert.True(t, strings.HasPrefix(buf.String(), "ran shared"))
|
||||||
|
}
|
||||||
+137
@@ -0,0 +1,137 @@
|
|||||||
|
package diag
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Exit statuses the client reports. They follow shell convention closely enough that a
|
||||||
|
// script can tell "you asked for something that does not exist" from "it ran and failed".
|
||||||
|
const (
|
||||||
|
// StatusOK means the command ran. Note that commands report their own user facing
|
||||||
|
// problems as prose and still exit 0, matching the ssh console.
|
||||||
|
StatusOK = 0
|
||||||
|
// StatusError means the command could not be completed.
|
||||||
|
StatusError = 1
|
||||||
|
// StatusUsage means the arguments were not valid for that command.
|
||||||
|
StatusUsage = 2
|
||||||
|
// StatusUnknownCommand means there is no such command.
|
||||||
|
StatusUnknownCommand = 127
|
||||||
|
)
|
||||||
|
|
||||||
|
// requestTimeout bounds how long a connected client may take to send its request line. There
|
||||||
|
// is deliberately no timeout on the response: `reload` runs every reload callback inline
|
||||||
|
// before it returns, and a slow one is not a reason to hang up on the operator.
|
||||||
|
const requestTimeout = 5 * time.Second
|
||||||
|
|
||||||
|
// Server serves a Registry over a stream listener. It knows nothing about unix sockets, so
|
||||||
|
// tests can drive it over a net.Pipe.
|
||||||
|
type Server struct {
|
||||||
|
l *slog.Logger
|
||||||
|
reg *Registry
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewServer(l *slog.Logger, reg *Registry) *Server {
|
||||||
|
return &Server{l: l, reg: reg}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Serve accepts connections until ln is closed. Cancelling ctx closes ln, which is what ends
|
||||||
|
// the accept loop; a listener closed underneath us is a normal shutdown, not an error.
|
||||||
|
func (s *Server) Serve(ctx context.Context, ln net.Listener) error {
|
||||||
|
go func() {
|
||||||
|
<-ctx.Done()
|
||||||
|
if err := ln.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
|
||||||
|
s.l.Warn("Failed to close the ctl listener", "error", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
for {
|
||||||
|
conn, err := ln.Accept()
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, net.ErrClosed) || ctx.Err() != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
go s.ServeConn(ctx, conn)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ServeConn handles one request and closes c.
|
||||||
|
func (s *Server) ServeConn(ctx context.Context, c net.Conn) {
|
||||||
|
defer func() {
|
||||||
|
if err := c.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
|
||||||
|
s.l.Debug("Failed to close a ctl connection", "error", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
if err := c.SetReadDeadline(time.Now().Add(requestTimeout)); err != nil {
|
||||||
|
s.l.Debug("Failed to set a ctl read deadline", "error", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
req, err := readRequest(bufio.NewReaderSize(c, maxRequest))
|
||||||
|
if err != nil {
|
||||||
|
s.l.Debug("Rejected a ctl request", "error", err)
|
||||||
|
// Best effort: the client may already be gone, and there is nowhere else to report it.
|
||||||
|
_ = writeEnd(c, StatusError, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// The request is in hand, so the command owns the rest of the connection's lifetime.
|
||||||
|
if err := c.SetReadDeadline(time.Time{}); err != nil {
|
||||||
|
s.l.Debug("Failed to clear the ctl read deadline", "error", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
s.l.Debug("Running a ctl command", "args", req.Args)
|
||||||
|
|
||||||
|
buf := bufio.NewWriterSize(&frameWriter{w: c}, outputBuffer)
|
||||||
|
dispatchErr := s.reg.DispatchArgs(req.Args, NewWriter(buf))
|
||||||
|
|
||||||
|
if err := buf.Flush(); err != nil {
|
||||||
|
s.l.Debug("Failed to flush ctl output", "error", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
status, msg := statusFor(dispatchErr)
|
||||||
|
if err := writeEnd(c, status, msg); err != nil {
|
||||||
|
s.l.Debug("Failed to write the ctl end frame", "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// StatusFor maps a dispatch error onto an exit status, for a transport that has somewhere to
|
||||||
|
// put one.
|
||||||
|
func StatusFor(err error) int {
|
||||||
|
status, _ := statusFor(err)
|
||||||
|
return status
|
||||||
|
}
|
||||||
|
|
||||||
|
// statusFor maps a dispatch error onto an exit status and, when the failure is ours to
|
||||||
|
// explain rather than one the command already wrote as prose, a message to go with it.
|
||||||
|
func statusFor(err error) (int, string) {
|
||||||
|
switch {
|
||||||
|
case err == nil:
|
||||||
|
return StatusOK, ""
|
||||||
|
case errors.Is(err, ErrUnknownCommand):
|
||||||
|
return StatusUnknownCommand, ""
|
||||||
|
case errors.Is(err, ErrUsage):
|
||||||
|
return StatusUsage, ""
|
||||||
|
default:
|
||||||
|
return StatusError, fmt.Sprintf("%s", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ErrNotSupported means this platform has no ctl transport. Windows is waiting on a named
|
||||||
|
// pipe implementation; mobile has no daemon for a CLI to attach to in the first place.
|
||||||
|
var ErrNotSupported = errors.New("nebula ctl is not supported on this platform")
|
||||||
|
|
||||||
|
// Listen creates the ctl listener at path. It is the platform boundary: everything above it
|
||||||
|
// in this package is portable.
|
||||||
|
func Listen(path string) (net.Listener, error) {
|
||||||
|
return listenSocket(path)
|
||||||
|
}
|
||||||
@@ -0,0 +1,273 @@
|
|||||||
|
//go:build !windows
|
||||||
|
|
||||||
|
package diag
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io/fs"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// testSocketPath returns a short socket path. t.TempDir on darwin lives under
|
||||||
|
// /var/folders/... and readily exceeds the 104 byte sun_path limit, which fails as a bare
|
||||||
|
// "invalid argument" a long way from the cause.
|
||||||
|
func testSocketPath(t *testing.T) string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
dir, err := os.MkdirTemp("/tmp", "nebctl")
|
||||||
|
require.NoError(t, err)
|
||||||
|
t.Cleanup(func() { _ = os.RemoveAll(dir) })
|
||||||
|
|
||||||
|
path := filepath.Join(dir, "sub", "ctl.sock")
|
||||||
|
require.LessOrEqual(t, len(path), maxSocketPath, "test socket path is too long for sun_path")
|
||||||
|
return path
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTestServer(t *testing.T) (*Registry, string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
reg := NewRegistry()
|
||||||
|
path := testSocketPath(t)
|
||||||
|
|
||||||
|
ln, err := Listen(path)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
srv := NewServer(slog.New(slog.DiscardHandler), reg)
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
assert.NoError(t, srv.Serve(ctx, ln))
|
||||||
|
}()
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
cancel()
|
||||||
|
wg.Wait()
|
||||||
|
})
|
||||||
|
|
||||||
|
return reg, path
|
||||||
|
}
|
||||||
|
|
||||||
|
func run(t *testing.T, path string, args ...string) (string, int, error) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
c, err := Dial(path)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer c.Close()
|
||||||
|
|
||||||
|
out := &bytes.Buffer{}
|
||||||
|
status, err := c.Run(args, out)
|
||||||
|
return out.String(), status, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServeConn(t *testing.T) {
|
||||||
|
t.Run("a command runs and its output comes back", func(t *testing.T) {
|
||||||
|
reg, path := newTestServer(t)
|
||||||
|
reg.RegisterCommand(&Command{
|
||||||
|
Name: "version",
|
||||||
|
ShortDescription: "prints a version",
|
||||||
|
Callback: func(fs any, a []string, w StringWriter) error {
|
||||||
|
return w.WriteLine("1.2.3")
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
out, status, err := run(t, path, "version")
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, StatusOK, status)
|
||||||
|
assert.Equal(t, "1.2.3\n", out)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("no args gets the command list", func(t *testing.T) {
|
||||||
|
_, path := newTestServer(t)
|
||||||
|
|
||||||
|
out, status, err := run(t, path)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, StatusOK, status)
|
||||||
|
assert.Contains(t, out, "Available commands:")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("an unknown command exits 127", func(t *testing.T) {
|
||||||
|
_, path := newTestServer(t)
|
||||||
|
|
||||||
|
out, status, err := run(t, path, "nope")
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, StatusUnknownCommand, status)
|
||||||
|
assert.Contains(t, out, "Did not understand: nope")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a bad flag exits 2", func(t *testing.T) {
|
||||||
|
reg, path := newTestServer(t)
|
||||||
|
var seen any
|
||||||
|
reg.RegisterCommand(testCommand("do-thing", &seen, nil))
|
||||||
|
|
||||||
|
out, status, err := run(t, path, "do-thing", "-nope")
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, StatusUsage, status)
|
||||||
|
assert.Contains(t, out, "flag provided but not defined")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a callback error exits 1 and reports why", func(t *testing.T) {
|
||||||
|
reg, path := newTestServer(t)
|
||||||
|
reg.RegisterCommand(&Command{
|
||||||
|
Name: "explode",
|
||||||
|
ShortDescription: "fails",
|
||||||
|
Callback: func(fs any, a []string, w StringWriter) error {
|
||||||
|
return errors.New("boom")
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
_, status, err := run(t, path, "explode")
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Equal(t, StatusError, status)
|
||||||
|
assert.Contains(t, err.Error(), "boom")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("output larger than the buffer arrives intact", func(t *testing.T) {
|
||||||
|
reg, path := newTestServer(t)
|
||||||
|
want := bytes.Repeat([]byte("x"), outputBuffer*3+7)
|
||||||
|
reg.RegisterCommand(&Command{
|
||||||
|
Name: "big",
|
||||||
|
ShortDescription: "writes a lot",
|
||||||
|
Callback: func(fs any, a []string, w StringWriter) error {
|
||||||
|
return w.WriteBytes(want)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
out, status, err := run(t, path, "big")
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, StatusOK, status)
|
||||||
|
assert.Equal(t, string(want), out)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("concurrent clients are all served", func(t *testing.T) {
|
||||||
|
reg, path := newTestServer(t)
|
||||||
|
reg.RegisterCommand(&Command{
|
||||||
|
Name: "slow",
|
||||||
|
ShortDescription: "takes a moment",
|
||||||
|
Callback: func(fs any, a []string, w StringWriter) error {
|
||||||
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
return w.WriteLine("done")
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for i := 0; i < 8; i++ {
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
out, status, err := run(t, path, "slow")
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, StatusOK, status)
|
||||||
|
assert.Equal(t, "done\n", out)
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a client that hangs up mid command does not take the server down", func(t *testing.T) {
|
||||||
|
reg, path := newTestServer(t)
|
||||||
|
reg.RegisterCommand(&Command{
|
||||||
|
Name: "version",
|
||||||
|
ShortDescription: "prints a version",
|
||||||
|
Callback: func(fs any, a []string, w StringWriter) error {
|
||||||
|
return w.WriteLine("1.2.3")
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
c, err := Dial(path)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, writeRequest(c.conn, []string{"version"}))
|
||||||
|
require.NoError(t, c.Close())
|
||||||
|
|
||||||
|
// The next client still gets served.
|
||||||
|
out, status, err := run(t, path, "version")
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, StatusOK, status)
|
||||||
|
assert.Equal(t, "1.2.3\n", out)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestListenSocket(t *testing.T) {
|
||||||
|
t.Run("the socket is 0600 inside a 0700 directory", func(t *testing.T) {
|
||||||
|
path := testSocketPath(t)
|
||||||
|
ln, err := Listen(path)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer ln.Close()
|
||||||
|
|
||||||
|
fi, err := os.Stat(path)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, os.FileMode(0600), fi.Mode().Perm(), "socket mode")
|
||||||
|
|
||||||
|
di, err := os.Stat(filepath.Dir(path))
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, os.FileMode(0700), di.Mode().Perm(), "socket directory mode")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("the socket is unlinked when the listener closes", func(t *testing.T) {
|
||||||
|
path := testSocketPath(t)
|
||||||
|
ln, err := Listen(path)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, ln.Close())
|
||||||
|
|
||||||
|
_, err = os.Stat(path)
|
||||||
|
assert.ErrorIs(t, err, fs.ErrNotExist)
|
||||||
|
})
|
||||||
|
|
||||||
|
// A crashed nebula leaves its socket behind, and the next one has to be able to start.
|
||||||
|
t.Run("a socket left behind by a dead nebula is replaced", func(t *testing.T) {
|
||||||
|
path := testSocketPath(t)
|
||||||
|
ln, err := Listen(path)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Close the listener without unlinking, the way a killed process leaves things.
|
||||||
|
unix, ok := ln.(*net.UnixListener)
|
||||||
|
require.True(t, ok)
|
||||||
|
unix.SetUnlinkOnClose(false)
|
||||||
|
require.NoError(t, ln.Close())
|
||||||
|
require.FileExists(t, path)
|
||||||
|
|
||||||
|
ln2, err := Listen(path)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.NoError(t, ln2.Close())
|
||||||
|
})
|
||||||
|
|
||||||
|
// Silently stealing it would break the nebula that got there first.
|
||||||
|
t.Run("a socket another nebula is serving is refused", func(t *testing.T) {
|
||||||
|
_, path := newTestServer(t)
|
||||||
|
|
||||||
|
_, err := Listen(path)
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "already being served")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a path that is not a socket is refused rather than removed", func(t *testing.T) {
|
||||||
|
path := testSocketPath(t)
|
||||||
|
require.NoError(t, os.MkdirAll(filepath.Dir(path), 0700))
|
||||||
|
require.NoError(t, os.WriteFile(path, []byte("precious"), 0600))
|
||||||
|
|
||||||
|
_, err := Listen(path)
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "is not a socket")
|
||||||
|
assert.FileExists(t, path, "the file must not have been removed")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a path too long for sun_path says so", func(t *testing.T) {
|
||||||
|
_, err := Listen("/tmp/" + fmt.Sprintf("%0*d", maxSocketPath, 0) + "/ctl.sock")
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "the maximum is")
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -0,0 +1,107 @@
|
|||||||
|
//go:build !windows
|
||||||
|
|
||||||
|
package diag
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io/fs"
|
||||||
|
"net"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"runtime"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// maxSocketPath is the smallest sun_path across the platforms nebula ships on: 104 bytes on
|
||||||
|
// darwin and the BSDs, 108 on Linux. Checking it ourselves turns a bare "invalid argument"
|
||||||
|
// into something an operator can act on.
|
||||||
|
const maxSocketPath = 103
|
||||||
|
|
||||||
|
// DefaultSocketPath is where nebula listens when ctl.socket is unset. An empty string means
|
||||||
|
// the platform has no sensible default and ctl stays off unless an operator names a path.
|
||||||
|
func DefaultSocketPath() string {
|
||||||
|
switch runtime.GOOS {
|
||||||
|
case "ios", "android":
|
||||||
|
// No daemon to attach to and no shell to attach from, and nowhere writable that
|
||||||
|
// would survive being guessed. Mobile embedders drive nebula through Control.
|
||||||
|
return ""
|
||||||
|
case "linux":
|
||||||
|
return "/run/nebula/ctl.sock"
|
||||||
|
default:
|
||||||
|
// /run does not exist on darwin, and /var/run is the portable spelling everywhere
|
||||||
|
// else nebula builds.
|
||||||
|
return "/var/run/nebula/ctl.sock"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// listenSocket creates the listening socket at path, taking over one a previous nebula left
|
||||||
|
// behind but refusing one that is still being served.
|
||||||
|
func listenSocket(path string) (net.Listener, error) {
|
||||||
|
if len(path) > maxSocketPath {
|
||||||
|
return nil, fmt.Errorf("socket path is %d bytes, the maximum is %d", len(path), maxSocketPath)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The directory, not the socket, is what enforces access control. net.Listen creates the
|
||||||
|
// socket with 0777&^umask, so with a typical 0022 umask it is world connectable for the
|
||||||
|
// window between bind and chmod. Nobody can traverse into a 0700 directory to reach it in
|
||||||
|
// that window, and unlike the socket's own mode, directory traversal is enforced
|
||||||
|
// consistently across every platform this file builds for.
|
||||||
|
dir := filepath.Dir(path)
|
||||||
|
if err := os.MkdirAll(dir, 0700); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create %s: %w", dir, err)
|
||||||
|
}
|
||||||
|
if err := os.Chmod(dir, 0700); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to set permissions on %s: %w", dir, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := clearStaleSocket(path); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
ln, err := net.Listen("unix", path)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Defence in depth behind the directory, for anyone who relocates the socket somewhere
|
||||||
|
// more permissive.
|
||||||
|
if err := os.Chmod(path, 0600); err != nil {
|
||||||
|
_ = ln.Close()
|
||||||
|
return nil, fmt.Errorf("failed to set permissions on %s: %w", path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return ln, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// dialSocket connects to a nebula serving at path.
|
||||||
|
func dialSocket(path string, timeout time.Duration) (net.Conn, error) {
|
||||||
|
return net.DialTimeout("unix", path, timeout)
|
||||||
|
}
|
||||||
|
|
||||||
|
// clearStaleSocket removes a socket a crashed nebula left behind, but refuses to steal one
|
||||||
|
// another nebula is still serving. Two instances on one host need two paths; they cannot
|
||||||
|
// share one, and silently taking the socket would break the instance that got there first.
|
||||||
|
func clearStaleSocket(path string) error {
|
||||||
|
fi, err := os.Lstat(path)
|
||||||
|
if errors.Is(err, fs.ErrNotExist) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if fi.Mode()&fs.ModeSocket == 0 {
|
||||||
|
return fmt.Errorf("%s exists and is not a socket, refusing to remove it", path)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A successful dial is the only reliable way to tell a live socket from an abandoned
|
||||||
|
// one; the inode looks identical either way.
|
||||||
|
c, err := net.DialTimeout("unix", path, 100*time.Millisecond)
|
||||||
|
if err == nil {
|
||||||
|
_ = c.Close()
|
||||||
|
return fmt.Errorf("%s is already being served, is another nebula running?", path)
|
||||||
|
}
|
||||||
|
|
||||||
|
return os.Remove(path)
|
||||||
|
}
|
||||||
@@ -0,0 +1,27 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package diag
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Windows has AF_UNIX since Windows 10 1803, but no way to secure the socket that resembles
|
||||||
|
// what the unix build does: os.Chmod cannot express an ACL, and a socket's reachability comes
|
||||||
|
// down to whatever its directory inherited. Doing this properly means a named pipe with an
|
||||||
|
// explicit security descriptor, which is a dependency and a design this change does not carry.
|
||||||
|
// Until then the stub keeps the package building and gives operators a real answer.
|
||||||
|
|
||||||
|
// DefaultSocketPath returns an empty string: there is no path worth defaulting to here.
|
||||||
|
func DefaultSocketPath() string {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func listenSocket(path string) (net.Listener, error) {
|
||||||
|
return nil, ErrNotSupported
|
||||||
|
}
|
||||||
|
|
||||||
|
func dialSocket(path string, timeout time.Duration) (net.Conn, error) {
|
||||||
|
return nil, ErrNotSupported
|
||||||
|
}
|
||||||
@@ -1,4 +1,4 @@
|
|||||||
package sshd
|
package diag
|
||||||
|
|
||||||
import "io"
|
import "io"
|
||||||
|
|
||||||
@@ -30,3 +30,9 @@ func (w *stringWriter) WriteBytes(b []byte) error {
|
|||||||
func (w *stringWriter) GetWriter() io.Writer {
|
func (w *stringWriter) GetWriter() io.Writer {
|
||||||
return w.w
|
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}
|
||||||
|
}
|
||||||
@@ -116,6 +116,9 @@ func newSimpleServerWithUdpAndUnsafeNetworks(v cert.Version, caCrt cert.Certific
|
|||||||
"key": string(myPrivKey),
|
"key": string(myPrivKey),
|
||||||
},
|
},
|
||||||
//"tun": m{"disabled": true},
|
//"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{
|
"firewall": m{
|
||||||
"outbound": []m{{
|
"outbound": []m{{
|
||||||
"proto": "any",
|
"proto": "any",
|
||||||
@@ -213,6 +216,9 @@ func newServer(caCrt []cert.Certificate, certs []cert.Certificate, key []byte, o
|
|||||||
"key": string(key),
|
"key": string(key),
|
||||||
},
|
},
|
||||||
//"tun": m{"disabled": true},
|
//"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{
|
"firewall": m{
|
||||||
"outbound": []m{{
|
"outbound": []m{{
|
||||||
"proto": "any",
|
"proto": "any",
|
||||||
|
|||||||
@@ -236,6 +236,30 @@ punchy:
|
|||||||
# Overriding this to "" is the same as "/" and will allow overwriting any path on the host.
|
# Overriding this to "" is the same as "/" and will allow overwriting any path on the host.
|
||||||
#sandbox_dir: /var/tmp/nebula-debug
|
#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.
|
# EXPERIMENTAL: relay support for networks that can't establish direct connections.
|
||||||
relay:
|
relay:
|
||||||
# Relays are a list of Nebula IP's that peers can use to relay packets to me.
|
# Relays are a list of Nebula IP's that peers can use to relay packets to me.
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ import (
|
|||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/cpupick"
|
"github.com/slackhq/nebula/cpupick"
|
||||||
|
"github.com/slackhq/nebula/diag"
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
"github.com/slackhq/nebula/sshd"
|
"github.com/slackhq/nebula/sshd"
|
||||||
@@ -68,7 +69,9 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
}
|
}
|
||||||
l.Info("Firewall started", "firewallHashes", fw.GetRuleHashes())
|
l.Info("Firewall started", "firewallHashes", fw.GetRuleHashes())
|
||||||
|
|
||||||
ssh, err := sshd.NewSSHServer(ctx, l.With("subsystem", "sshd"))
|
commands := diag.NewRegistry()
|
||||||
|
|
||||||
|
ssh, err := sshd.NewSSHServer(ctx, l.With("subsystem", "sshd"), commands)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, util.ContextualizeIfNeeded("Error while creating SSH server", err)
|
return nil, util.ContextualizeIfNeeded("Error while creating SSH server", err)
|
||||||
}
|
}
|
||||||
@@ -322,13 +325,20 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
return nil, util.ContextualizeIfNeeded("Failed to start stats emitter", err)
|
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 {
|
if configTest {
|
||||||
return nil, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
go ifce.emitStats(ctx, c.GetDuration("stats.interval", time.Second*10))
|
go ifce.emitStats(ctx, c.GetDuration("stats.interval", time.Second*10))
|
||||||
|
|
||||||
attachCommands(l, c, ssh, ifce)
|
attachCommands(l, c, commands, ifce)
|
||||||
|
|
||||||
networkChanges := udp.NewNetworkChangeMonitor(ctx, l, c)
|
networkChanges := udp.NewNetworkChangeMonitor(ctx, l, c)
|
||||||
|
|
||||||
@@ -339,6 +349,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
ctx: ctx,
|
ctx: ctx,
|
||||||
cancel: cancel,
|
cancel: cancel,
|
||||||
sshStart: sshStart,
|
sshStart: sshStart,
|
||||||
|
ctlStart: ctlServer.Start,
|
||||||
statsStart: stats.Start,
|
statsStart: stats.Start,
|
||||||
dnsStart: ds.Start,
|
dnsStart: ds.Start,
|
||||||
lighthouseStart: lightHouse.StartUpdateWorker,
|
lighthouseStart: lightHouse.StartUpdateWorker,
|
||||||
|
|||||||
@@ -1,62 +1,19 @@
|
|||||||
package nebula
|
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 (
|
import (
|
||||||
"bytes"
|
|
||||||
"encoding/json"
|
|
||||||
"errors"
|
|
||||||
"flag"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"maps"
|
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
|
||||||
"runtime"
|
|
||||||
"runtime/pprof"
|
|
||||||
"sort"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/header"
|
|
||||||
"github.com/slackhq/nebula/logging"
|
|
||||||
"github.com/slackhq/nebula/sshd"
|
"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) {
|
func wireSSHReload(l *slog.Logger, ssh *sshd.SSHServer, c *config.C) {
|
||||||
c.RegisterReloadCallback(func(c *config.C) {
|
c.RegisterReloadCallback(func(c *config.C) {
|
||||||
if c.GetBool("sshd.enabled", false) {
|
if c.GetBool("sshd.enabled", false) {
|
||||||
@@ -197,862 +154,3 @@ func configSSH(l *slog.Logger, ssh *sshd.SSHServer, c *config.C) (func(), error)
|
|||||||
|
|
||||||
return runner, nil
|
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
|
|
||||||
}
|
|
||||||
|
|||||||
+6
-19
@@ -9,8 +9,9 @@ import (
|
|||||||
"net"
|
"net"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"github.com/armon/go-radix"
|
|
||||||
"golang.org/x/crypto/ssh"
|
"golang.org/x/crypto/ssh"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/diag"
|
||||||
)
|
)
|
||||||
|
|
||||||
type SSHServer struct {
|
type SSHServer struct {
|
||||||
@@ -25,9 +26,8 @@ type SSHServer struct {
|
|||||||
trustedKeys map[string]map[string]bool
|
trustedKeys map[string]map[string]bool
|
||||||
trustedCAs []ssh.PublicKey
|
trustedCAs []ssh.PublicKey
|
||||||
|
|
||||||
// List of available commands
|
// The commands this server serves. Shared with every other transport, see diag.Registry.
|
||||||
helpCommand *Command
|
commands *diag.Registry
|
||||||
commands *radix.Tree
|
|
||||||
listener net.Listener
|
listener net.Listener
|
||||||
|
|
||||||
// ctx parents per-Run contexts. Cancelling it (e.g. via Control.Stop) tears the server down even
|
// ctx parents per-Run contexts. Cancelling it (e.g. via Control.Stop) tears the server down even
|
||||||
@@ -38,11 +38,11 @@ type SSHServer struct {
|
|||||||
// NewSSHServer creates a new ssh server rigged with default commands and prepares to listen.
|
// 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
|
// 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.
|
// (e.g. on Control.Stop) tears down active sessions and closes the listener.
|
||||||
func NewSSHServer(ctx context.Context, l *slog.Logger) (*SSHServer, error) {
|
func NewSSHServer(ctx context.Context, l *slog.Logger, commands *diag.Registry) (*SSHServer, error) {
|
||||||
s := &SSHServer{
|
s := &SSHServer{
|
||||||
trustedKeys: make(map[string]map[string]bool),
|
trustedKeys: make(map[string]map[string]bool),
|
||||||
l: l,
|
l: l,
|
||||||
commands: radix.New(),
|
commands: commands,
|
||||||
ctx: ctx,
|
ctx: ctx,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -90,14 +90,6 @@ func NewSSHServer(ctx context.Context, l *slog.Logger) (*SSHServer, error) {
|
|||||||
ServerVersion: fmt.Sprintf("SSH-2.0-Nebula???"),
|
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
|
return s, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -160,11 +152,6 @@ func (s *SSHServer) AddAuthorizedKey(user, pubKey string) error {
|
|||||||
return nil
|
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
|
// 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
|
// 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.
|
// rather than carrying a permanently-cancelled context across runs.
|
||||||
|
|||||||
+18
-45
@@ -1,37 +1,38 @@
|
|||||||
package sshd
|
package sshd
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"sort"
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/anmitsu/go-shlex"
|
|
||||||
"github.com/armon/go-radix"
|
|
||||||
"golang.org/x/crypto/ssh"
|
"golang.org/x/crypto/ssh"
|
||||||
"golang.org/x/term"
|
"golang.org/x/term"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/diag"
|
||||||
)
|
)
|
||||||
|
|
||||||
type session struct {
|
type session struct {
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
c *ssh.ServerConn
|
c *ssh.ServerConn
|
||||||
term *term.Terminal
|
term *term.Terminal
|
||||||
commands *radix.Tree
|
commands *diag.Registry
|
||||||
cancel func()
|
cancel func()
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewSession(commands *radix.Tree, conn *ssh.ServerConn, chans <-chan ssh.NewChannel, cancel func(), l *slog.Logger) *session {
|
func NewSession(commands *diag.Registry, conn *ssh.ServerConn, chans <-chan ssh.NewChannel, cancel func(), l *slog.Logger) *session {
|
||||||
s := &session{
|
s := &session{
|
||||||
commands: radix.NewFromMap(commands.ToMap()),
|
// A copy, so the logout command this session adds for itself stays invisible to every
|
||||||
|
// other session and to `nebula ctl`.
|
||||||
|
commands: commands.Clone(),
|
||||||
l: l,
|
l: l,
|
||||||
c: conn,
|
c: conn,
|
||||||
cancel: cancel,
|
cancel: cancel,
|
||||||
}
|
}
|
||||||
|
|
||||||
s.commands.Insert("logout", &Command{
|
s.commands.RegisterCommand(&diag.Command{
|
||||||
Name: "logout",
|
Name: "logout",
|
||||||
ShortDescription: "Ends the current session",
|
ShortDescription: "Ends the current session",
|
||||||
Callback: func(a any, args []string, w StringWriter) error {
|
Callback: func(a any, args []string, w diag.StringWriter) error {
|
||||||
s.Close()
|
s.Close()
|
||||||
return nil
|
return nil
|
||||||
},
|
},
|
||||||
@@ -87,9 +88,11 @@ func (s *session) handleRequests(in <-chan *ssh.Request, channel ssh.Channel) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
req.Reply(true, nil)
|
req.Reply(true, nil)
|
||||||
s.dispatchCommand(payload.Value, &stringWriter{channel})
|
dErr := s.commands.Dispatch(payload.Value, diag.NewWriter(channel))
|
||||||
|
|
||||||
status := struct{ Status uint32 }{uint32(0)}
|
// Report a real exit status rather than a hardcoded zero, so that
|
||||||
|
// `ssh nebula-host list-hostmap` is scriptable the same way `nebula ctl` is.
|
||||||
|
status := struct{ Status uint32 }{uint32(diag.StatusFor(dErr))}
|
||||||
channel.SendRequest("exit-status", false, ssh.Marshal(status))
|
channel.SendRequest("exit-status", false, ssh.Marshal(status))
|
||||||
channel.Close()
|
channel.Close()
|
||||||
return
|
return
|
||||||
@@ -111,7 +114,7 @@ func (s *session) createTerm(channel ssh.Channel) *term.Terminal {
|
|||||||
term.AutoCompleteCallback = func(line string, pos int, key rune) (newLine string, newPos int, ok bool) {
|
term.AutoCompleteCallback = func(line string, pos int, key rune) (newLine string, newPos int, ok bool) {
|
||||||
// key 9 is tab
|
// key 9 is tab
|
||||||
if key == 9 {
|
if key == 9 {
|
||||||
cmds := matchCommand(s.commands, line)
|
cmds := s.commands.Match(line)
|
||||||
if len(cmds) == 1 {
|
if len(cmds) == 1 {
|
||||||
return cmds[0] + " ", len(cmds[0]) + 1, true
|
return cmds[0] + " ", len(cmds[0]) + 1, true
|
||||||
}
|
}
|
||||||
@@ -128,49 +131,19 @@ func (s *session) createTerm(channel ssh.Channel) *term.Terminal {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *session) handleInput() {
|
func (s *session) handleInput() {
|
||||||
w := &stringWriter{w: s.term}
|
w := diag.NewWriter(s.term)
|
||||||
for {
|
for {
|
||||||
line, err := s.term.ReadLine()
|
line, err := s.term.ReadLine()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
s.dispatchCommand(line, w)
|
// The interactive console reports problems on the terminal the user is already
|
||||||
|
// looking at, so the error is nothing extra to say here.
|
||||||
|
_ = s.commands.Dispatch(line, w)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *session) dispatchCommand(line string, w StringWriter) {
|
|
||||||
args, err := shlex.Split(line, true)
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(args) == 0 {
|
|
||||||
dumpCommands(s.commands, w)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
c, err := lookupCommand(s.commands, args[0])
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if c == nil {
|
|
||||||
err := w.WriteLine(fmt.Sprintf("did not understand: %s", line))
|
|
||||||
_ = err
|
|
||||||
|
|
||||||
dumpCommands(s.commands, w)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if checkHelpArgs(args) {
|
|
||||||
s.dispatchCommand(fmt.Sprintf("%s %s", "help", c.Name), w)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
_ = execCommand(c, args[1:], w)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *session) Close() {
|
func (s *session) Close() {
|
||||||
s.c.Close()
|
s.c.Close()
|
||||||
s.cancel()
|
s.cancel()
|
||||||
|
|||||||
Reference in New Issue
Block a user