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