Add nebula ctl, a local socket for the debug commands

Every diagnostic command nebula has was reachable through exactly one door:
the built-in ssh debug server. That server is off by default, and turning it
on means generating a host key, writing an sshd block with authorized public
keys, and SIGHUPing the daemon. That is a lot of ceremony to answer "what
version is this node running".

Nebula now serves the same commands over a local unix socket, enabled by
default, and `nebula ctl <command>` runs them. The socket lives in a 0700
directory so filesystem permissions are the access control; no keys, nothing
on the network. Failing to create it is logged and never blocks startup.

The command registry was already transport neutral, so this is mostly new
transport rather than new commands:

  - diag/ holds the registry, dispatch, writer and wire protocol, moved out
    of sshd because none of it was ever about ssh. sshd and ctl.go dispatch
    against one shared registry.
  - commands.go holds every command implementation, moved out of ssh.go
    (which was 85% not ssh) and renamed off the ssh prefix. Adding a command
    there makes it available over both transports.
  - ssh.go keeps only host keys, authorized users, and the listen address.
  - ctl.go supervises the socket, following the statsServer lifecycle shape.

The wire protocol frames the response rather than terminating it, because
print-cert -raw and list-hostmap -json both emit arbitrary bytes that no
sentinel could safely delimit. argv travels as a list so quoting survives.
Exit statuses are real: 0, 2 for usage, 127 for an unknown command.

Two things fall out. The ssh console now reports a real exit status instead
of a hardcoded zero, so `ssh host list-hostmap` is scriptable too. And eight
command callbacks that silently returned nil on a flags type mismatch now
report it, which the exit status makes visible.

Windows is a stub returning a clear "not supported" until it gets a named
pipe with a security descriptor; iOS and Android are never enabled, having no
daemon for a CLI to attach to.

Breaking for embedders of the sshd package: NewSSHServer takes a
*diag.Registry, SSHServer.RegisterCommand is gone in favor of registering on
that registry, and the command types live in diag rather than sshd.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014fya5fTXGiwX72FUmoL9y3
This commit is contained in:
Matt Richardson
2026-09-09 17:07:35 -04:00
co-authored by Claude Opus 5
parent 89178f45ba
commit 14a4d87faf
26 changed files with 3113 additions and 976 deletions
+41
View File
@@ -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()
}
+174
View File
@@ -0,0 +1,174 @@
package diag
import (
"errors"
"flag"
"fmt"
"sort"
"strings"
"github.com/armon/go-radix"
)
var (
// ErrUnknownCommand is returned by the Registry when the first argument names no
// registered command. The user has already been told so on their writer.
ErrUnknownCommand = errors.New("unknown command")
// ErrUsage wraps a flag parsing failure. The flag package has already written the
// details to the caller's writer by the time this is returned, so a transport should
// use it only to pick an exit status.
ErrUsage = errors.New("usage")
)
// CommandFlags is a function called before help or command execution to parse command line flags
// It should return a flag.FlagSet instance and a pointer to the struct that will contain parsed flags
type CommandFlags func() (*flag.FlagSet, any)
// CommandCallback is the function called when your command should execute.
// fs will be a a pointer to the struct provided by Command.Flags callback, if there was one. -h and -help are reserved
// and handled automatically for you.
// a will be any unconsumed arguments, if no Command.Flags was available this will be all the flags passed in.
// w is the writer to use when sending messages back to the client.
// If an error is returned by the callback it is logged locally, the callback should handle messaging errors to the user
// where appropriate
type CommandCallback func(fs any, a []string, w StringWriter) error
type Command struct {
Name string
ShortDescription string
Help string
Flags CommandFlags
Callback CommandCallback
}
func execCommand(c *Command, args []string, w StringWriter) error {
var (
fl *flag.FlagSet
fs any
)
if c.Flags != nil {
fl, fs = c.Flags()
if fl != nil {
// SetOutput() here in case fl.Parse dumps usage.
fl.SetOutput(w.GetWriter())
err := fl.Parse(args)
if err != nil {
// fl.Parse has dumped error information to the user via the w writer, so
// the wrapper exists purely so a transport can tell a usage problem from a
// command that ran and failed.
return fmt.Errorf("%w: %w", ErrUsage, err)
}
args = fl.Args()
}
}
return c.Callback(fs, args, w)
}
func dumpCommands(c *radix.Tree, w StringWriter) {
err := w.WriteLine("Available commands:")
if err != nil {
return
}
cmds := make([]string, 0)
for _, l := range allCommands(c) {
cmds = append(cmds, fmt.Sprintf("%s - %s", l.Name, l.ShortDescription))
}
sort.Strings(cmds)
_ = w.Write(strings.Join(cmds, "\n") + "\n\n")
}
func lookupCommand(c *radix.Tree, sCmd string) (*Command, error) {
cmd, ok := c.Get(sCmd)
if !ok {
return nil, nil
}
command, ok := cmd.(*Command)
if !ok {
return nil, errors.New("failed to cast command")
}
return command, nil
}
func matchCommand(c *radix.Tree, cmd string) []string {
cmds := make([]string, 0)
c.WalkPrefix(cmd, func(found string, v any) bool {
cmds = append(cmds, found)
return false
})
sort.Strings(cmds)
return cmds
}
func allCommands(c *radix.Tree) []*Command {
cmds := make([]*Command, 0)
c.WalkPrefix("", func(found string, v any) bool {
cmd, ok := v.(*Command)
if ok {
cmds = append(cmds, cmd)
}
return false
})
return cmds
}
func helpCallback(commands *radix.Tree, a []string, w StringWriter) (err error) {
// Just typed help
if len(a) == 0 {
dumpCommands(commands, w)
return nil
}
// We are printing a specific commands help text
cmd, err := lookupCommand(commands, a[0])
if err != nil {
return
}
if cmd != nil {
err = w.WriteLine(fmt.Sprintf("%s - %s", cmd.Name, cmd.ShortDescription))
if err != nil {
return err
}
if cmd.Help != "" {
err = w.WriteLine(fmt.Sprintf(" %s", cmd.Help))
if err != nil {
return err
}
}
if cmd.Flags != nil {
fs, _ := cmd.Flags()
if fs != nil {
fs.SetOutput(w.GetWriter())
fs.PrintDefaults()
}
}
return nil
}
err = w.WriteLine("Command not available " + a[0])
if err != nil {
return err
}
return nil
}
func checkHelpArgs(args []string) bool {
for _, a := range args {
if a == "-h" || a == "-help" {
return true
}
}
return false
}
+226
View File
@@ -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])
}
}
}
+144
View File
@@ -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")
})
}
+125
View File
@@ -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)
}
+167
View File
@@ -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
View File
@@ -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)
}
+273
View File
@@ -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")
})
}
+107
View File
@@ -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)
}
+27
View File
@@ -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
}
+38
View File
@@ -0,0 +1,38 @@
package diag
import "io"
type StringWriter interface {
WriteLine(string) error
Write(string) error
WriteBytes([]byte) error
GetWriter() io.Writer
}
type stringWriter struct {
w io.Writer
}
func (w *stringWriter) WriteLine(s string) error {
return w.Write(s + "\n")
}
func (w *stringWriter) Write(s string) error {
_, err := w.w.Write([]byte(s))
return err
}
func (w *stringWriter) WriteBytes(b []byte) error {
_, err := w.w.Write(b)
return err
}
func (w *stringWriter) GetWriter() io.Writer {
return w.w
}
// NewWriter adapts an io.Writer to the StringWriter commands are handed. Transports
// implement their own framing behind w; the commands never know the difference.
func NewWriter(w io.Writer) StringWriter {
return &stringWriter{w: w}
}