diff --git a/CHANGELOG.md b/CHANGELOG.md index ed089185..9634cc7a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,28 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- New `nebula ctl ` 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 ` rather than always reporting + success, so commands run that way are scriptable. +- The debug and administrative commands moved out of `ssh.go` into `commands.go` and are no longer tied to + ssh: both the ssh console and `nebula ctl` dispatch against one shared registry, so a command added in + one place is available over both. Embedders of the `sshd` package are affected: `sshd.NewSSHServer` now + takes a `*diag.Registry`, `sshd.SSHServer.RegisterCommand` is gone in favor of registering on that + registry directly, and the command types now live in the `diag` package rather than being re-exported + from `sshd`. + ## [1.11.1] - 2026-08-21 See the [v1.11.1](https://github.com/slackhq/nebula/milestone/30?closed=1) milestone for a complete list of changes. diff --git a/cmd/nebula/ctl.go b/cmd/nebula/ctl.go new file mode 100644 index 00000000..71f51f0a --- /dev/null +++ b/cmd/nebula/ctl.go @@ -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 [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] [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] +} diff --git a/cmd/nebula/ctl_test.go b/cmd/nebula/ctl_test.go new file mode 100644 index 00000000..a339b506 --- /dev/null +++ b/cmd/nebula/ctl_test.go @@ -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)) +} diff --git a/cmd/nebula/main.go b/cmd/nebula/main.go index 3c786b84..21e1be4c 100644 --- a/cmd/nebula/main.go +++ b/cmd/nebula/main.go @@ -32,11 +32,26 @@ func init() { } func main() { + // Subcommands are dispatched before flag.Parse, because flag.Parse stops at the first + // non-flag argument and everything after `ctl` has to reach the running nebula's own flag + // parser untouched. Nothing here looks at -json or a vpn address. + if len(os.Args) > 1 && os.Args[1] == "ctl" { + os.Exit(ctlMain(os.Args[2:])) + } + configPath := flag.String("config", "", "Path to either a file or directory to load configuration from") configTest := flag.Bool("test", false, "Test the config and print the end result. Non zero exit indicates a faulty config") printVersion := flag.Bool("version", false, "Print version") printUsage := flag.Bool("help", false, "Print command line usage") + flag.Usage = func() { + out := flag.CommandLine.Output() + fmt.Fprintf(out, "Usage of %s:\n", os.Args[0]) + flag.PrintDefaults() + fmt.Fprintf(out, "\nCommands:\n") + fmt.Fprintf(out, " ctl [command]\n\tRun a debug command against the running nebula on this host.\n\tRun `nebula ctl` on its own for the list of commands.\n") + } + flag.Parse() if *printVersion { diff --git a/commands.go b/commands.go new file mode 100644 index 00000000..b16d9fbb --- /dev/null +++ b/commands.go @@ -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 +} diff --git a/commands_test.go b/commands_test.go new file mode 100644 index 00000000..5f667cd8 --- /dev/null +++ b/commands_test.go @@ -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) + } + }) +} diff --git a/control.go b/control.go index 41bc4f97..61793a6b 100644 --- a/control.go +++ b/control.go @@ -50,6 +50,7 @@ type Control struct { ctx context.Context cancel context.CancelFunc sshStart func() + ctlStart func() statsStart func() dnsStart func() lighthouseStart func() @@ -99,6 +100,9 @@ func (c *Control) Start() error { if c.sshStart != nil { go c.sshStart() } + if c.ctlStart != nil { + go c.ctlStart() + } if c.statsStart != nil { go c.statsStart() } diff --git a/ctl.go b/ctl.go new file mode 100644 index 00000000..45df03fb --- /dev/null +++ b/ctl.go @@ -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) + } +} diff --git a/ctl_test.go b/ctl_test.go new file mode 100644 index 00000000..185723f5 --- /dev/null +++ b/ctl_test.go @@ -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) } diff --git a/diag/client.go b/diag/client.go new file mode 100644 index 00000000..ca96ac7b --- /dev/null +++ b/diag/client.go @@ -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() +} diff --git a/sshd/command.go b/diag/command.go similarity index 84% rename from sshd/command.go rename to diag/command.go index 7323d120..3bb2fd2f 100644 --- a/sshd/command.go +++ b/diag/command.go @@ -1,4 +1,4 @@ -package sshd +package diag import ( "errors" @@ -10,6 +10,17 @@ import ( "github.com/armon/go-radix" ) +var ( + // ErrUnknownCommand is returned by the Registry when the first argument names no + // registered command. The user has already been told so on their writer. + ErrUnknownCommand = errors.New("unknown command") + + // ErrUsage wraps a flag parsing failure. The flag package has already written the + // details to the caller's writer by the time this is returned, so a transport should + // use it only to pick an exit status. + ErrUsage = errors.New("usage") +) + // CommandFlags is a function called before help or command execution to parse command line flags // It should return a flag.FlagSet instance and a pointer to the struct that will contain parsed flags type CommandFlags func() (*flag.FlagSet, any) @@ -44,8 +55,10 @@ func execCommand(c *Command, args []string, w StringWriter) error { fl.SetOutput(w.GetWriter()) err := fl.Parse(args) if err != nil { - // fl.Parse has dumped error information to the user via the w writer. - return err + // fl.Parse has dumped error information to the user via the w writer, so + // the wrapper exists purely so a transport can tell a usage problem from a + // command that ran and failed. + return fmt.Errorf("%w: %w", ErrUsage, err) } args = fl.Args() } diff --git a/diag/proto.go b/diag/proto.go new file mode 100644 index 00000000..e426745c --- /dev/null +++ b/diag/proto.go @@ -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]) + } + } +} diff --git a/diag/proto_test.go b/diag/proto_test.go new file mode 100644 index 00000000..40858f98 --- /dev/null +++ b/diag/proto_test.go @@ -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") + }) +} diff --git a/diag/registry.go b/diag/registry.go new file mode 100644 index 00000000..919e46fd --- /dev/null +++ b/diag/registry.go @@ -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 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) +} diff --git a/diag/registry_test.go b/diag/registry_test.go new file mode 100644 index 00000000..db4a646d --- /dev/null +++ b/diag/registry_test.go @@ -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")) +} diff --git a/diag/server.go b/diag/server.go new file mode 100644 index 00000000..d848ca87 --- /dev/null +++ b/diag/server.go @@ -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) +} diff --git a/diag/server_test.go b/diag/server_test.go new file mode 100644 index 00000000..c8866c53 --- /dev/null +++ b/diag/server_test.go @@ -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") + }) +} diff --git a/diag/socket_unix.go b/diag/socket_unix.go new file mode 100644 index 00000000..a5b8bad4 --- /dev/null +++ b/diag/socket_unix.go @@ -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) +} diff --git a/diag/socket_windows.go b/diag/socket_windows.go new file mode 100644 index 00000000..0cfe6d1e --- /dev/null +++ b/diag/socket_windows.go @@ -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 +} diff --git a/sshd/writer.go b/diag/writer.go similarity index 66% rename from sshd/writer.go rename to diag/writer.go index 8354c094..990cc765 100644 --- a/sshd/writer.go +++ b/diag/writer.go @@ -1,4 +1,4 @@ -package sshd +package diag import "io" @@ -30,3 +30,9 @@ func (w *stringWriter) WriteBytes(b []byte) error { func (w *stringWriter) GetWriter() io.Writer { return w.w } + +// NewWriter adapts an io.Writer to the StringWriter commands are handed. Transports +// implement their own framing behind w; the commands never know the difference. +func NewWriter(w io.Writer) StringWriter { + return &stringWriter{w: w} +} diff --git a/e2e/helpers_test.go b/e2e/helpers_test.go index 1691aeab..1f3720c1 100644 --- a/e2e/helpers_test.go +++ b/e2e/helpers_test.go @@ -116,6 +116,9 @@ func newSimpleServerWithUdpAndUnsafeNetworks(v cert.Version, caCrt cert.Certific "key": string(myPrivKey), }, //"tun": m{"disabled": true}, + // Several tests bring up more than one nebula in this process, and they would all + // contend for the same default ctl socket path. None of them exercise it. + "ctl": m{"enabled": false}, "firewall": m{ "outbound": []m{{ "proto": "any", @@ -213,6 +216,9 @@ func newServer(caCrt []cert.Certificate, certs []cert.Certificate, key []byte, o "key": string(key), }, //"tun": m{"disabled": true}, + // Several tests bring up more than one nebula in this process, and they would all + // contend for the same default ctl socket path. None of them exercise it. + "ctl": m{"enabled": false}, "firewall": m{ "outbound": []m{{ "proto": "any", diff --git a/examples/config.yml b/examples/config.yml index 20a85883..ffd898f3 100644 --- a/examples/config.yml +++ b/examples/config.yml @@ -236,6 +236,30 @@ punchy: # Overriding this to "" is the same as "/" and will allow overwriting any path on the host. #sandbox_dir: /var/tmp/nebula-debug +# ctl exposes nebula's debug and administrative commands over a local unix socket, so that `nebula ctl ` can +# reach the same commands the sshd block offers above without running an ssh server. Run `nebula ctl` on its own for the +# list of commands. Anyone who can open the socket can do everything the ssh console can, including closing tunnels, +# changing remotes, and writing profile data to disk, so the socket lives in a directory only the user nebula runs as +# can enter. Enabled by default. Not supported on Windows yet, and never enabled on iOS or Android. +#ctl: + # Toggles the feature. This setting is reloadable. + #enabled: true + + # socket is the unix socket to listen on. The parent directory is created if it is missing and made readable only by + # the user nebula runs as, and a socket left behind by a crashed nebula is replaced. Defaults to /run/nebula/ctl.sock + # on Linux and /var/run/nebula/ctl.sock everywhere else; running nebula as a non-root user means picking a path it can + # write. Two nebulas on one host need two paths, the second to start will log that the socket is already being served + # and carry on without one. `nebula ctl` reads this value from the same config file when it is given -config, and + # otherwise assumes the default above. This setting is reloadable. + #socket: /run/nebula/ctl.sock + + # sandbox_dir restricts the file paths the profiling commands (start-cpu-profile, save-heap-profile, + # save-mutex-profile) may write, exactly like sshd.sandbox_dir above, which it defaults to. Note that these paths are + # resolved by the nebula process and not by the shell running `nebula ctl`, so a relative path lands in this directory + # rather than in your working directory, and under a systemd unit with PrivateTmp=yes it lands somewhere your shell + # cannot see at all. The directory is NOT automatically created. + #sandbox_dir: /var/tmp/nebula-debug + # EXPERIMENTAL: relay support for networks that can't establish direct connections. relay: # Relays are a list of Nebula IP's that peers can use to relay packets to me. diff --git a/main.go b/main.go index ee52e7b7..fedd1c2e 100644 --- a/main.go +++ b/main.go @@ -14,6 +14,7 @@ import ( "github.com/slackhq/nebula/config" "github.com/slackhq/nebula/cpupick" + "github.com/slackhq/nebula/diag" "github.com/slackhq/nebula/noiseutil" "github.com/slackhq/nebula/overlay" "github.com/slackhq/nebula/sshd" @@ -68,7 +69,9 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev } l.Info("Firewall started", "firewallHashes", fw.GetRuleHashes()) - ssh, err := sshd.NewSSHServer(ctx, l.With("subsystem", "sshd")) + commands := diag.NewRegistry() + + ssh, err := sshd.NewSSHServer(ctx, l.With("subsystem", "sshd"), commands) if err != nil { return nil, util.ContextualizeIfNeeded("Error while creating SSH server", err) } @@ -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) } + // Built before the configTest return so that a bad ctl block fails `nebula -test`. It only + // holds the registry, which attachCommands populates below, and reads nothing until Start. + ctlServer, err := newCtlServerFromConfig(ctx, l.With("subsystem", "ctl"), c, commands) + if err != nil { + return nil, util.ContextualizeIfNeeded("Failed to configure the ctl socket", err) + } + if configTest { return nil, nil } go ifce.emitStats(ctx, c.GetDuration("stats.interval", time.Second*10)) - attachCommands(l, c, ssh, ifce) + attachCommands(l, c, commands, ifce) networkChanges := udp.NewNetworkChangeMonitor(ctx, l, c) @@ -339,6 +349,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev ctx: ctx, cancel: cancel, sshStart: sshStart, + ctlStart: ctlServer.Start, statsStart: stats.Start, dnsStart: ds.Start, lighthouseStart: lightHouse.StartUpdateWorker, diff --git a/ssh.go b/ssh.go index 3863b5ec..42ae66b0 100644 --- a/ssh.go +++ b/ssh.go @@ -1,62 +1,19 @@ package nebula +// Configuration and lifecycle for the ssh debug console. The commands it serves are not +// defined here; see commands.go, which registers them for every transport. + import ( - "bytes" - "encoding/json" - "errors" - "flag" "fmt" "log/slog" - "maps" "net" - "net/netip" "os" - "path/filepath" - "runtime" - "runtime/pprof" - "sort" - "strconv" "strings" "github.com/slackhq/nebula/config" - "github.com/slackhq/nebula/header" - "github.com/slackhq/nebula/logging" "github.com/slackhq/nebula/sshd" ) -type sshListHostMapFlags struct { - Json bool - Pretty bool - ByIndex bool -} - -type sshPrintCertFlags struct { - Json bool - Pretty bool - Raw bool -} - -type sshPrintTunnelFlags struct { - Pretty bool -} - -type sshChangeRemoteFlags struct { - Address string -} - -type sshCloseTunnelFlags struct { - LocalOnly bool -} - -type sshCreateTunnelFlags struct { - Address string -} - -type sshDeviceInfoFlags struct { - Json bool - Pretty bool -} - func wireSSHReload(l *slog.Logger, ssh *sshd.SSHServer, c *config.C) { c.RegisterReloadCallback(func(c *config.C) { if c.GetBool("sshd.enabled", false) { @@ -197,862 +154,3 @@ func configSSH(l *slog.Logger, ssh *sshd.SSHServer, c *config.C) (func(), error) return runner, nil } - -func attachCommands(l *slog.Logger, c *config.C, ssh *sshd.SSHServer, f *Interface) { - // sandboxDir defaults to a dir in temp. The intention is that end user will - // create this dir as needed. Overriding this config value to "" allows - // writing to anywhere in the system. - defaultDir := filepath.Join(os.TempDir(), "nebula-debug") - sandboxDir := c.GetString("sshd.sandbox_dir", defaultDir) - - ssh.RegisterCommand(&sshd.Command{ - Name: "list-hostmap", - ShortDescription: "List all known previously connected hosts", - Flags: func() (*flag.FlagSet, any) { - fl := flag.NewFlagSet("", flag.ContinueOnError) - s := sshListHostMapFlags{} - fl.BoolVar(&s.Json, "json", false, "outputs as json with more information") - fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json") - fl.BoolVar(&s.ByIndex, "by-index", false, "gets all hosts in the hostmap from the index table") - return fl, &s - }, - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshListHostMap(f.hostMap, fs, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "list-pending-hostmap", - ShortDescription: "List all handshaking hosts", - Flags: func() (*flag.FlagSet, any) { - fl := flag.NewFlagSet("", flag.ContinueOnError) - s := sshListHostMapFlags{} - fl.BoolVar(&s.Json, "json", false, "outputs as json with more information") - fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json") - fl.BoolVar(&s.ByIndex, "by-index", false, "gets all hosts in the hostmap from the index table") - return fl, &s - }, - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshListHostMap(f.handshakeManager, fs, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "list-lighthouse-addrmap", - ShortDescription: "List all lighthouse map entries", - Flags: func() (*flag.FlagSet, any) { - fl := flag.NewFlagSet("", flag.ContinueOnError) - s := sshListHostMapFlags{} - fl.BoolVar(&s.Json, "json", false, "outputs as json with more information") - fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json") - return fl, &s - }, - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshListLighthouseMap(f.lightHouse, fs, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "reload", - ShortDescription: "Reloads configuration from disk, same as sending HUP to the process", - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshReload(c, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "start-cpu-profile", - ShortDescription: "Starts a cpu profile and write output to the provided file, ex: `cpu-profile.pb.gz`", - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshStartCpuProfile(sandboxDir, fs, a, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "stop-cpu-profile", - ShortDescription: "Stops a cpu profile and writes output to the previously provided file", - Callback: func(fs any, a []string, w sshd.StringWriter) error { - pprof.StopCPUProfile() - return w.WriteLine("If a CPU profile was running it is now stopped") - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "save-heap-profile", - ShortDescription: "Saves a heap profile to the provided path, ex: `heap-profile.pb.gz`", - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshGetHeapProfile(sandboxDir, fs, a, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "mutex-profile-fraction", - ShortDescription: "Gets or sets runtime.SetMutexProfileFraction", - Callback: sshMutexProfileFraction, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "save-mutex-profile", - ShortDescription: "Saves a mutex profile to the provided path, ex: `mutex-profile.pb.gz`", - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshGetMutexProfile(sandboxDir, fs, a, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "log-level", - ShortDescription: "Gets or sets the current log level", - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshLogLevel(l, fs, a, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "log-format", - ShortDescription: "Gets or sets the current log format", - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshLogFormat(l, fs, a, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "version", - ShortDescription: "Prints the currently running version of nebula", - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshVersion(f, fs, a, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "device-info", - ShortDescription: "Prints information about the network device.", - Flags: func() (*flag.FlagSet, any) { - fl := flag.NewFlagSet("", flag.ContinueOnError) - s := sshDeviceInfoFlags{} - fl.BoolVar(&s.Json, "json", false, "outputs as json with more information") - fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json") - return fl, &s - }, - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshDeviceInfo(f, fs, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "print-cert", - ShortDescription: "Prints the current certificate being used or the certificate for the provided vpn addr", - Flags: func() (*flag.FlagSet, any) { - fl := flag.NewFlagSet("", flag.ContinueOnError) - s := sshPrintCertFlags{} - fl.BoolVar(&s.Json, "json", false, "outputs as json") - fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json") - fl.BoolVar(&s.Raw, "raw", false, "raw prints the PEM encoded certificate, not compatible with -json or -pretty") - return fl, &s - }, - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshPrintCert(f, fs, a, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "print-tunnel", - ShortDescription: "Prints json details about a tunnel for the provided vpn addr", - Flags: func() (*flag.FlagSet, any) { - fl := flag.NewFlagSet("", flag.ContinueOnError) - s := sshPrintTunnelFlags{} - fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json") - return fl, &s - }, - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshPrintTunnel(f, fs, a, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "print-relays", - ShortDescription: "Prints json details about all relay info", - Flags: func() (*flag.FlagSet, any) { - fl := flag.NewFlagSet("", flag.ContinueOnError) - s := sshPrintTunnelFlags{} - fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json") - return fl, &s - }, - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshPrintRelays(f, fs, a, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "change-remote", - ShortDescription: "Changes the remote address used in the tunnel for the provided vpn addr", - Flags: func() (*flag.FlagSet, any) { - fl := flag.NewFlagSet("", flag.ContinueOnError) - s := sshChangeRemoteFlags{} - fl.StringVar(&s.Address, "address", "", "The new remote address, ip:port") - return fl, &s - }, - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshChangeRemote(f, fs, a, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "close-tunnel", - ShortDescription: "Closes a tunnel for the provided vpn addr", - Flags: func() (*flag.FlagSet, any) { - fl := flag.NewFlagSet("", flag.ContinueOnError) - s := sshCloseTunnelFlags{} - fl.BoolVar(&s.LocalOnly, "local-only", false, "Disables notifying the remote that the tunnel is shutting down") - return fl, &s - }, - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshCloseTunnel(f, fs, a, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "create-tunnel", - ShortDescription: "Creates a tunnel for the provided vpn address", - Help: "The lighthouses will be queried for real addresses but you can provide one as well.", - Flags: func() (*flag.FlagSet, any) { - fl := flag.NewFlagSet("", flag.ContinueOnError) - s := sshCreateTunnelFlags{} - fl.StringVar(&s.Address, "address", "", "Optionally provide a real remote address, ip:port ") - return fl, &s - }, - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshCreateTunnel(f, fs, a, w) - }, - }) - - ssh.RegisterCommand(&sshd.Command{ - Name: "query-lighthouse", - ShortDescription: "Query the lighthouses for the provided vpn address", - Help: "This command is asynchronous. Only currently known udp addresses will be printed.", - Callback: func(fs any, a []string, w sshd.StringWriter) error { - return sshQueryLighthouse(f, fs, a, w) - }, - }) -} - -func sshListHostMap(hl controlHostLister, a any, w sshd.StringWriter) error { - fs, ok := a.(*sshListHostMapFlags) - if !ok { - return nil - } - - var hm []ControlHostInfo - if fs.ByIndex { - hm = listHostMapIndexes(hl) - } else { - hm = listHostMapHosts(hl) - } - - sort.Slice(hm, func(i, j int) bool { - return hm[i].VpnAddrs[0].Compare(hm[j].VpnAddrs[0]) < 0 - }) - - if fs.Json || fs.Pretty { - js := json.NewEncoder(w.GetWriter()) - if fs.Pretty { - js.SetIndent("", " ") - } - - err := js.Encode(hm) - if err != nil { - return nil - } - - } else { - for _, v := range hm { - err := w.WriteLine(fmt.Sprintf("%s: %s", v.VpnAddrs, v.RemoteAddrs)) - if err != nil { - return err - } - } - } - - return nil -} - -func sshListLighthouseMap(lightHouse *LightHouse, a any, w sshd.StringWriter) error { - fs, ok := a.(*sshListHostMapFlags) - if !ok { - return nil - } - - type lighthouseInfo struct { - VpnAddr string `json:"vpnAddr"` - Addrs *CacheMap `json:"addrs"` - } - - lightHouse.RLock() - addrMap := make([]lighthouseInfo, len(lightHouse.addrMap)) - x := 0 - for k, v := range lightHouse.addrMap { - addrMap[x] = lighthouseInfo{ - VpnAddr: k.String(), - Addrs: v.CopyCache(), - } - x++ - } - lightHouse.RUnlock() - - sort.Slice(addrMap, func(i, j int) bool { - return strings.Compare(addrMap[i].VpnAddr, addrMap[j].VpnAddr) < 0 - }) - - if fs.Json || fs.Pretty { - js := json.NewEncoder(w.GetWriter()) - if fs.Pretty { - js.SetIndent("", " ") - } - - err := js.Encode(addrMap) - if err != nil { - return nil - } - - } else { - for _, v := range addrMap { - b, err := json.Marshal(v.Addrs) - if err != nil { - return err - } - err = w.WriteLine(fmt.Sprintf("%s: %s", v.VpnAddr, string(b))) - if err != nil { - return err - } - } - } - - return nil -} - -// sshSanitizeFilePath validates that the given file path is within the sandbox directory. -// If sandboxDir is empty, the path is returned as-is for backwards compatibility. -func sshSanitizeFilePath(sandboxDir, filePath string) (string, error) { - if sandboxDir == "" { - return filePath, nil - } - - // Clean and resolve the path relative to the sandbox directory - if !filepath.IsAbs(filePath) { - filePath = filepath.Join(sandboxDir, filePath) - } - cleaned := filepath.Clean(filePath) - - // Ensure the resolved path is within the sandbox directory - cleanedSandbox := filepath.Clean(sandboxDir) - if cleaned == cleanedSandbox { - return "", fmt.Errorf("path %q resolves to the sandbox directory itself %q", filePath, sandboxDir) - } - if !strings.HasPrefix(cleaned, cleanedSandbox+string(filepath.Separator)) { - return "", fmt.Errorf("path %q is outside the sandbox directory %q", filePath, sandboxDir) - } - - return cleaned, nil -} - -func sshStartCpuProfile(sandboxDir string, fs any, a []string, w sshd.StringWriter) error { - if len(a) == 0 { - err := w.WriteLine("No path to write profile provided") - return err - } - - filePath, err := sshSanitizeFilePath(sandboxDir, a[0]) - if err != nil { - return w.WriteLine(err.Error()) - } - - file, err := os.Create(filePath) - if err != nil { - err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err)) - return err - } - - err = pprof.StartCPUProfile(file) - if err != nil { - err = w.WriteLine(fmt.Sprintf("Unable to start cpu profile: %s", err)) - return err - } - - err = w.WriteLine(fmt.Sprintf("Started cpu profile, issue stop-cpu-profile to write the output to %s", a)) - return err -} - -func sshVersion(ifce *Interface, fs any, a []string, w sshd.StringWriter) error { - return w.WriteLine(fmt.Sprintf("%s", ifce.version)) -} - -func sshQueryLighthouse(ifce *Interface, fs any, a []string, w sshd.StringWriter) error { - if len(a) == 0 { - return w.WriteLine("No vpn address was provided") - } - - vpnAddr, err := netip.ParseAddr(a[0]) - if err != nil { - return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0])) - } - - if !vpnAddr.IsValid() { - return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0])) - } - - var cm *CacheMap - rl := ifce.lightHouse.Query(vpnAddr) - if rl != nil { - cm = rl.CopyCache() - } - return json.NewEncoder(w.GetWriter()).Encode(cm) -} - -func sshCloseTunnel(ifce *Interface, fs any, a []string, w sshd.StringWriter) error { - flags, ok := fs.(*sshCloseTunnelFlags) - if !ok { - return nil - } - - if len(a) == 0 { - return w.WriteLine("No vpn address was provided") - } - - vpnAddr, err := netip.ParseAddr(a[0]) - if err != nil { - return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0])) - } - - if !vpnAddr.IsValid() { - return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0])) - } - - hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr) - if hostInfo == nil { - return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn address: %v", a[0])) - } - - if !flags.LocalOnly { - ifce.send( - header.CloseTunnel, - 0, - hostInfo.ConnectionState, - hostInfo, - []byte{}, - make([]byte, 12, 12), - make([]byte, mtu), - ) - } - - ifce.closeTunnel(hostInfo) - return w.WriteLine("Closed") -} - -func sshCreateTunnel(ifce *Interface, fs any, a []string, w sshd.StringWriter) error { - flags, ok := fs.(*sshCreateTunnelFlags) - if !ok { - return nil - } - - if len(a) == 0 { - return w.WriteLine("No vpn address was provided") - } - - vpnAddr, err := netip.ParseAddr(a[0]) - if err != nil { - return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0])) - } - - if !vpnAddr.IsValid() { - return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0])) - } - - hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr) - if hostInfo != nil { - return w.WriteLine(fmt.Sprintf("Tunnel already exists")) - } - - hostInfo = ifce.handshakeManager.QueryVpnAddr(vpnAddr) - if hostInfo != nil { - return w.WriteLine(fmt.Sprintf("Tunnel already handshaking")) - } - - var addr netip.AddrPort - if flags.Address != "" { - addr, err = netip.ParseAddrPort(flags.Address) - if err != nil { - return w.WriteLine("Address could not be parsed") - } - } - - hostInfo = ifce.handshakeManager.StartHandshake(vpnAddr, nil) - if addr.IsValid() { - hostInfo.SetRemote(addr) - } - - return w.WriteLine("Created") -} - -func sshChangeRemote(ifce *Interface, fs any, a []string, w sshd.StringWriter) error { - flags, ok := fs.(*sshChangeRemoteFlags) - if !ok { - return nil - } - - if len(a) == 0 { - return w.WriteLine("No vpn address was provided") - } - - if flags.Address == "" { - return w.WriteLine("No address was provided") - } - - addr, err := netip.ParseAddrPort(flags.Address) - if err != nil { - return w.WriteLine("Address could not be parsed") - } - - vpnAddr, err := netip.ParseAddr(a[0]) - if err != nil { - return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0])) - } - - if !vpnAddr.IsValid() { - return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0])) - } - - hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr) - if hostInfo == nil { - return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn address: %v", a[0])) - } - - hostInfo.SetRemote(addr) - return w.WriteLine("Changed") -} - -func sshGetHeapProfile(sandboxDir string, fs any, a []string, w sshd.StringWriter) error { - if len(a) == 0 { - return w.WriteLine("No path to write profile provided") - } - - filePath, err := sshSanitizeFilePath(sandboxDir, a[0]) - if err != nil { - return w.WriteLine(err.Error()) - } - - file, err := os.Create(filePath) - if err != nil { - err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err)) - return err - } - - err = pprof.WriteHeapProfile(file) - if err != nil { - err = w.WriteLine(fmt.Sprintf("Unable to write profile: %s", err)) - return err - } - - err = w.WriteLine(fmt.Sprintf("Mem profile created at %s", a)) - return err -} - -func sshMutexProfileFraction(fs any, a []string, w sshd.StringWriter) error { - if len(a) == 0 { - rate := runtime.SetMutexProfileFraction(-1) - return w.WriteLine(fmt.Sprintf("Current value: %d", rate)) - } - - newRate, err := strconv.Atoi(a[0]) - if err != nil { - return w.WriteLine(fmt.Sprintf("Invalid argument: %s", a[0])) - } - - oldRate := runtime.SetMutexProfileFraction(newRate) - return w.WriteLine(fmt.Sprintf("New value: %d. Old value: %d", newRate, oldRate)) -} - -func sshGetMutexProfile(sandboxDir string, fs any, a []string, w sshd.StringWriter) error { - if len(a) == 0 { - return w.WriteLine("No path to write profile provided") - } - - filePath, err := sshSanitizeFilePath(sandboxDir, a[0]) - if err != nil { - return w.WriteLine(err.Error()) - } - - file, err := os.Create(filePath) - if err != nil { - return w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err)) - } - defer file.Close() - - mutexProfile := pprof.Lookup("mutex") - if mutexProfile == nil { - return w.WriteLine("Unable to get pprof.Lookup(\"mutex\")") - } - - err = mutexProfile.WriteTo(file, 0) - if err != nil { - return w.WriteLine(fmt.Sprintf("Unable to write profile: %s", err)) - } - - return w.WriteLine(fmt.Sprintf("Mutex profile created at %s", a)) -} - -func sshLogLevel(l *slog.Logger, fs any, a []string, w sshd.StringWriter) error { - ctrl, ok := l.Handler().(interface { - GetLevel() slog.Level - SetLevel(slog.Level) - }) - if !ok { - return w.WriteLine("Log level is not reconfigurable on this logger") - } - - if len(a) == 0 { - return w.WriteLine(fmt.Sprintf("Log level is: %s", logging.LevelName(ctrl.GetLevel()))) - } - - level, err := logging.ParseLevel(strings.ToLower(a[0])) - if err != nil { - return w.WriteLine(fmt.Sprintf("Unknown log level %s. Possible log levels: trace, debug, info, warn, error", a)) - } - - ctrl.SetLevel(level) - return w.WriteLine(fmt.Sprintf("Log level is: %s", logging.LevelName(ctrl.GetLevel()))) -} - -func sshLogFormat(l *slog.Logger, fs any, a []string, w sshd.StringWriter) error { - ctrl, ok := l.Handler().(interface { - GetFormat() string - SetFormat(string) error - }) - if !ok { - return w.WriteLine("Log format is not reconfigurable on this logger") - } - - if len(a) == 0 { - return w.WriteLine(fmt.Sprintf("Log format is: %s", ctrl.GetFormat())) - } - - if err := ctrl.SetFormat(strings.ToLower(a[0])); err != nil { - return err - } - return w.WriteLine(fmt.Sprintf("Log format is: %s", ctrl.GetFormat())) -} - -func sshPrintCert(ifce *Interface, fs any, a []string, w sshd.StringWriter) error { - args, ok := fs.(*sshPrintCertFlags) - if !ok { - return nil - } - - cert := ifce.pki.getCertState().GetDefaultCertificate() - if len(a) > 0 { - vpnAddr, err := netip.ParseAddr(a[0]) - if err != nil { - return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0])) - } - - if !vpnAddr.IsValid() { - return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0])) - } - - hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr) - if hostInfo == nil { - return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn addr: %v", a[0])) - } - - cert = hostInfo.GetCert().Certificate - } - - if args.Json || args.Pretty { - b, err := cert.MarshalJSON() - if err != nil { - return nil - } - - if args.Pretty { - buf := new(bytes.Buffer) - err := json.Indent(buf, b, "", " ") - b = buf.Bytes() - if err != nil { - return nil - } - } - - return w.WriteBytes(b) - } - - if args.Raw { - b, err := cert.MarshalPEM() - if err != nil { - return nil - } - - return w.WriteBytes(b) - } - - return w.WriteLine(cert.String()) -} - -func sshPrintRelays(ifce *Interface, fs any, a []string, w sshd.StringWriter) error { - args, ok := fs.(*sshPrintTunnelFlags) - if !ok { - w.WriteLine(fmt.Sprintf("sshPrintRelays failed to convert args type")) - return nil - } - - relays := map[uint32]*HostInfo{} - ifce.hostMap.Lock() - maps.Copy(relays, ifce.hostMap.Relays) - ifce.hostMap.Unlock() - - type RelayFor struct { - Error error - Type string - State string - PeerAddr netip.Addr - LocalIndex uint32 - RemoteIndex uint32 - RelayedThrough []netip.Addr - } - - type RelayOutput struct { - NebulaAddr netip.Addr - RelayForAddrs []RelayFor - } - - type CmdOutput struct { - Relays []*RelayOutput - } - - co := CmdOutput{} - - enc := json.NewEncoder(w.GetWriter()) - - if args.Pretty { - enc.SetIndent("", " ") - } - - for k, v := range relays { - ro := RelayOutput{NebulaAddr: v.vpnAddrs[0]} - co.Relays = append(co.Relays, &ro) - relayHI := ifce.hostMap.QueryVpnAddr(v.vpnAddrs[0]) - if relayHI == nil { - ro.RelayForAddrs = append(ro.RelayForAddrs, RelayFor{Error: errors.New("could not find hostinfo")}) - continue - } - for _, vpnAddr := range relayHI.relayState.CopyRelayForIps() { - rf := RelayFor{Error: nil} - r, ok := relayHI.relayState.GetRelayForByAddr(vpnAddr) - if ok { - t := "" - switch r.Type { - case ForwardingType: - t = "forwarding" - case TerminalType: - t = "terminal" - default: - t = "unknown" - } - - s := "" - switch r.State { - case Requested: - s = "requested" - case Established: - s = "established" - default: - s = "unknown" - } - - rf.LocalIndex = r.LocalIndex - rf.RemoteIndex = r.RemoteIndex - rf.PeerAddr = r.PeerAddr - rf.Type = t - rf.State = s - if rf.LocalIndex != k { - rf.Error = fmt.Errorf("hostmap LocalIndex '%v' does not match RelayState LocalIndex", k) - } - } - relayedHI := ifce.hostMap.QueryVpnAddr(vpnAddr) - if relayedHI != nil { - rf.RelayedThrough = append(rf.RelayedThrough, relayedHI.relayState.CopyRelayIps()...) - } - - ro.RelayForAddrs = append(ro.RelayForAddrs, rf) - } - } - err := enc.Encode(co) - if err != nil { - return err - } - return nil -} - -func sshPrintTunnel(ifce *Interface, fs any, a []string, w sshd.StringWriter) error { - args, ok := fs.(*sshPrintTunnelFlags) - if !ok { - return nil - } - - if len(a) == 0 { - return w.WriteLine("No vpn address was provided") - } - - vpnAddr, err := netip.ParseAddr(a[0]) - if err != nil { - return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0])) - } - - if !vpnAddr.IsValid() { - return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0])) - } - - hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr) - if hostInfo == nil { - return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn addr: %v", a[0])) - } - - enc := json.NewEncoder(w.GetWriter()) - if args.Pretty { - enc.SetIndent("", " ") - } - - return enc.Encode(copyHostInfo(hostInfo, ifce.hostMap.GetPreferredRanges())) -} - -func sshDeviceInfo(ifce *Interface, fs any, w sshd.StringWriter) error { - - data := struct { - Name string `json:"name"` - Cidr []netip.Prefix `json:"cidr"` - }{ - Name: ifce.inside.Name(), - Cidr: make([]netip.Prefix, len(ifce.inside.Networks())), - } - - copy(data.Cidr, ifce.inside.Networks()) - - flags, ok := fs.(*sshDeviceInfoFlags) - if !ok { - return fmt.Errorf("internal error: expected flags to be sshDeviceInfoFlags but was %+v", fs) - } - - if flags.Json || flags.Pretty { - js := json.NewEncoder(w.GetWriter()) - if flags.Pretty { - js.SetIndent("", " ") - } - - return js.Encode(data) - } else { - return w.WriteLine(fmt.Sprintf("name=%v cidr=%v", data.Name, data.Cidr)) - } -} - -func sshReload(c *config.C, w sshd.StringWriter) error { - err := w.WriteLine("Reloading config") - c.ReloadConfig() - return err -} diff --git a/sshd/server.go b/sshd/server.go index e0ee9364..a17aac99 100644 --- a/sshd/server.go +++ b/sshd/server.go @@ -9,8 +9,9 @@ import ( "net" "sync" - "github.com/armon/go-radix" "golang.org/x/crypto/ssh" + + "github.com/slackhq/nebula/diag" ) type SSHServer struct { @@ -25,10 +26,9 @@ type SSHServer struct { trustedKeys map[string]map[string]bool trustedCAs []ssh.PublicKey - // List of available commands - helpCommand *Command - commands *radix.Tree - listener net.Listener + // The commands this server serves. Shared with every other transport, see diag.Registry. + commands *diag.Registry + listener net.Listener // ctx parents per-Run contexts. Cancelling it (e.g. via Control.Stop) tears the server down even // across reloads, since each Run derives a fresh child rather than reusing this one directly. @@ -38,11 +38,11 @@ type SSHServer struct { // NewSSHServer creates a new ssh server rigged with default commands and prepares to listen. // The ssh server's context is parented off the supplied ctx so cancelling it // (e.g. on Control.Stop) tears down active sessions and closes the listener. -func NewSSHServer(ctx context.Context, l *slog.Logger) (*SSHServer, error) { +func NewSSHServer(ctx context.Context, l *slog.Logger, commands *diag.Registry) (*SSHServer, error) { s := &SSHServer{ trustedKeys: make(map[string]map[string]bool), l: l, - commands: radix.New(), + commands: commands, ctx: ctx, } @@ -90,14 +90,6 @@ func NewSSHServer(ctx context.Context, l *slog.Logger) (*SSHServer, error) { ServerVersion: fmt.Sprintf("SSH-2.0-Nebula???"), } - s.RegisterCommand(&Command{ - Name: "help", - ShortDescription: "prints available commands or help for specific usage info", - Callback: func(a any, args []string, w StringWriter) error { - return helpCallback(s.commands, args, w) - }, - }) - return s, nil } @@ -160,11 +152,6 @@ func (s *SSHServer) AddAuthorizedKey(user, pubKey string) error { return nil } -// RegisterCommand adds a command that can be run by a user, by default only `help` is available -func (s *SSHServer) RegisterCommand(c *Command) { - s.commands.Insert(c.Name, c) -} - // Run begins listening and accepting connections. Each invocation derives a fresh per-Run context // from the constructor-supplied ctx so a Stop+Run sequence (used by config reload) starts clean // rather than carrying a permanently-cancelled context across runs. diff --git a/sshd/session.go b/sshd/session.go index 1c8e1a9b..8172bb7b 100644 --- a/sshd/session.go +++ b/sshd/session.go @@ -1,37 +1,38 @@ package sshd import ( - "fmt" "log/slog" "sort" "strings" - "github.com/anmitsu/go-shlex" - "github.com/armon/go-radix" "golang.org/x/crypto/ssh" "golang.org/x/term" + + "github.com/slackhq/nebula/diag" ) type session struct { l *slog.Logger c *ssh.ServerConn term *term.Terminal - commands *radix.Tree + commands *diag.Registry cancel func() } -func NewSession(commands *radix.Tree, conn *ssh.ServerConn, chans <-chan ssh.NewChannel, cancel func(), l *slog.Logger) *session { +func NewSession(commands *diag.Registry, conn *ssh.ServerConn, chans <-chan ssh.NewChannel, cancel func(), l *slog.Logger) *session { s := &session{ - commands: radix.NewFromMap(commands.ToMap()), + // A copy, so the logout command this session adds for itself stays invisible to every + // other session and to `nebula ctl`. + commands: commands.Clone(), l: l, c: conn, cancel: cancel, } - s.commands.Insert("logout", &Command{ + s.commands.RegisterCommand(&diag.Command{ Name: "logout", ShortDescription: "Ends the current session", - Callback: func(a any, args []string, w StringWriter) error { + Callback: func(a any, args []string, w diag.StringWriter) error { s.Close() return nil }, @@ -87,9 +88,11 @@ func (s *session) handleRequests(in <-chan *ssh.Request, channel ssh.Channel) { } req.Reply(true, nil) - s.dispatchCommand(payload.Value, &stringWriter{channel}) + dErr := s.commands.Dispatch(payload.Value, diag.NewWriter(channel)) - status := struct{ Status uint32 }{uint32(0)} + // Report a real exit status rather than a hardcoded zero, so that + // `ssh nebula-host list-hostmap` is scriptable the same way `nebula ctl` is. + status := struct{ Status uint32 }{uint32(diag.StatusFor(dErr))} channel.SendRequest("exit-status", false, ssh.Marshal(status)) channel.Close() return @@ -111,7 +114,7 @@ func (s *session) createTerm(channel ssh.Channel) *term.Terminal { term.AutoCompleteCallback = func(line string, pos int, key rune) (newLine string, newPos int, ok bool) { // key 9 is tab if key == 9 { - cmds := matchCommand(s.commands, line) + cmds := s.commands.Match(line) if len(cmds) == 1 { return cmds[0] + " ", len(cmds[0]) + 1, true } @@ -128,49 +131,19 @@ func (s *session) createTerm(channel ssh.Channel) *term.Terminal { } func (s *session) handleInput() { - w := &stringWriter{w: s.term} + w := diag.NewWriter(s.term) for { line, err := s.term.ReadLine() if err != nil { break } - s.dispatchCommand(line, w) + // The interactive console reports problems on the terminal the user is already + // looking at, so the error is nothing extra to say here. + _ = s.commands.Dispatch(line, w) } } -func (s *session) dispatchCommand(line string, w StringWriter) { - args, err := shlex.Split(line, true) - if err != nil { - return - } - - if len(args) == 0 { - dumpCommands(s.commands, w) - return - } - - c, err := lookupCommand(s.commands, args[0]) - if err != nil { - return - } - - if c == nil { - err := w.WriteLine(fmt.Sprintf("did not understand: %s", line)) - _ = err - - dumpCommands(s.commands, w) - return - } - - if checkHelpArgs(args) { - s.dispatchCommand(fmt.Sprintf("%s %s", "help", c.Name), w) - return - } - - _ = execCommand(c, args[1:], w) -} - func (s *session) Close() { s.c.Close() s.cancel()