diff --git a/.golangci.yml b/.golangci.yml index 5f366c072..2f31946f8 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -49,11 +49,13 @@ linters: - $gostd - github.com/presmihaylov/shard/models - github.com/presmihaylov/shard/pkg/pty + - github.com/presmihaylov/shard/pkg/term - github.com/presmihaylov/shard/pkg/vzshim - github.com/presmihaylov/shard/services/client - github.com/presmihaylov/shard/services/sandbox - github.com/presmihaylov/shard/services/daemon - github.com/presmihaylov/shard/services/serve + - github.com/presmihaylov/shard/services/setup # models/ is a leaf so that it never participates in an import cycle. models-is-a-leaf: files: diff --git a/AGENTS.md b/AGENTS.md index d4221ff3a..d3c282729 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -74,6 +74,7 @@ pkg/reflink/ one clone by reference, and whether a directory's fil pkg/xfs/ mkfs.xfs, the loop mount and the fstab line of one image pkg/vz/ the Virtualization.framework driver: the shim protocol, its client and its server pkg/vzshim/ the shim binary embedded in the daemon, installed and ad-hoc signed on first use +pkg/term/ the terminal prompts and the live checklist shard setup draws services/sandbox/ the orchestrator: the lifecycle verbs the daemon serves services/image/ pull, unpack, cache policy @@ -87,6 +88,7 @@ services/daemon/ shard daemon: the wiring of every layer, and the back services/api/ the REST handlers the daemon serves over its unix socket services/client/ the typed client of that API, which the thin CLI verbs call services/serve/ the TCP front: a bearer token, and the bytes onto that socket +services/setup/ shard setup: inspect this host, install a provider and the daemon, or save a remote services/provider/gvisor/ implements models.Provider on gVisor services/provider/sysbox/ implements models.Provider on Sysbox services/provider/runc/ implements models.Provider on bare runc @@ -111,9 +113,10 @@ website/ useshards.com: Astro + Starlight, landing at /, docs a driver and it belongs in `services/`. `depguard` enforces this in CI. - **Dependencies point one way: `cli` to `services` to `pkg`.** `models` sits under all of them. -- **`cli/` imports `services/client`, `pkg/pty`, `pkg/vzshim`, `models`, the request - types in `services/sandbox`, and `services/daemon` and `services/serve` for the - two verbs that are a process rather than a client. Nothing else.** A verb holds no +- **`cli/` imports `services/client`, `pkg/pty`, `pkg/term`, `pkg/vzshim`, `models`, the + request types in `services/sandbox`, and `services/daemon`, `services/serve` and + `services/setup` for the three verbs that are a process rather than a client. + Nothing else.** A verb holds no store and no provider: it asks the socket. `depguard` enforces the allow list in CI. - **`models/` is one package with several files, and it is a leaf.** It imports diff --git a/README.md b/README.md index a1827c2c7..6e3047628 100644 --- a/README.md +++ b/README.md @@ -154,6 +154,7 @@ are in those directories, and [the release guide](docs/release.md) says how an S ## Documentation +- [Set up a host or a remote connection](docs/setup.md) - [CLI commands and options](docs/cli.md) - [Daemon, REST API, and remote access](docs/daemon.md) - [Provider capabilities and limits](docs/provider.md) diff --git a/cli/cli.go b/cli/cli.go index 60ee587f6..d7ac17598 100644 --- a/cli/cli.go +++ b/cli/cli.go @@ -2,6 +2,7 @@ package cli import ( + "cmp" "context" "errors" "flag" @@ -54,6 +55,9 @@ type App struct { // plainWarned is shared by every copy of the App one run makes, so a verb that builds two clients warns once. plainWarned *sync.Once + + // remoteFlag says --remote was passed, so even an empty one names the target and the saved connection names none. + remoteFlag bool } // stdin is what exec hands the guest and what secret set reads the value from. @@ -254,6 +258,7 @@ func commands() []command { {name: "daemon", run: App.daemon, subs: []command{{name: "status", run: App.daemonStatus}}}, {name: "info", run: App.info}, {name: "serve", run: App.serve}, + {name: "setup", run: App.setup}, {name: "tokens", subs: []command{ {name: "mint", run: App.tokensMint}, {name: "list", aliases: []string{"ls"}, run: App.tokensList}, @@ -388,6 +393,7 @@ func (a *App) parseGlobals(args []string) ([]string, error) { if err := parseVerb(flags, args); err != nil { return nil, err } + flags.Visit(func(f *flag.Flag) { a.remoteFlag = a.remoteFlag || f.Name == "remote" }) // --version answers before the root is checked, so it never fails. if showVersion { @@ -435,11 +441,15 @@ func (h *hostList) Set(value string) error { return nil } -// client speaks to the daemon on the socket, or through --remote; a verb asks only after its flags parsed, so --help reads no token. +// client speaks to the daemon on the socket, or through --remote, SHARD_REMOTE or the saved connection; a verb asks only after its flags parsed, so --help reads no token. func (a App) client() (*client.Client, error) { - // The key and the certificate are read here, so a bad one fails before the verb dials. - if a.Remote != "" { - c, err := client.NewRemoteFromEnv(a.Remote) + saved, err := a.saved() + if err != nil { + return nil, err + } + // a.Remote is already --remote or SHARD_REMOTE (fromEnv), so the saved remote comes last; a bad key or certificate fails before the dial. + if remote := cmp.Or(a.Remote, saved.Remote); remote != "" { + c, err := client.NewRemoteFromEnv(remote, saved) if err != nil { return nil, err } @@ -469,6 +479,22 @@ func (a App) localClient(verb string) (*client.Client, error) { // hostOnly refuses a remote for a verb that acts on this host, so it never reports the local result as the server's. func (a App) hostOnly(verb string) error { + if err := a.noRemote(verb); err != nil { + return err + } + saved, err := a.saved() + if err != nil { + return err + } + if saved.Remote == "" { + return nil + } + + return fmt.Errorf("shard %s runs on the daemon host only and cannot reach the %v; remove it with shard setup to run it here", verb, saved) +} + +// noRemote is hostOnly for daemon and serve: the saved connection names where commands go, never where a daemon runs. +func (a App) noRemote(verb string) error { if a.Remote == "" { return nil } @@ -476,6 +502,19 @@ func (a App) hostOnly(verb string) error { return fmt.Errorf("shard %s runs on the daemon host only and cannot reach %s; unset --remote and %s to run it there", verb, a.Remote, client.RemoteEnv) } +// saved is the connection shard setup saved; an explicit --remote "" asks for the socket, so it reads none. +func (a App) saved() (client.Config, error) { + if a.remoteFlag && a.Remote == "" { + return client.Config{}, nil + } + path, err := client.ConfigPath(os.Getenv) + if err != nil { + return client.Config{}, err + } + + return client.LoadConfig(path) +} + // gotArgs echoes what a verb refused, quoted, so the error shows what was typed rather than a count. func gotArgs(args []string) string { if len(args) == 0 { diff --git a/cli/config_test.go b/cli/config_test.go new file mode 100644 index 000000000..1708dbe0a --- /dev/null +++ b/cli/config_test.go @@ -0,0 +1,257 @@ +package cli + +import ( + "bytes" + "context" + "fmt" + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "slices" + "strings" + "sync" + "testing" + "time" + + "github.com/presmihaylov/shard/services/client" +) + +// leakKey is a synthetic key no real credential holds, so any text that carries it is a leak. +const leakKey = "shard657-synthetic-key-5a4b3c2d" + +// isolateConfig points the saved connection at an empty directory of its own, so no test reads or writes the user's. +func isolateConfig() (func() error, error) { + dir, err := os.MkdirTemp("", "shard-config") + if err != nil { + return nil, fmt.Errorf("make a configuration directory: %w", err) + } + if err := os.Setenv(client.ConfigHomeEnv, dir); err != nil { + return nil, fmt.Errorf("set %s: %w", client.ConfigHomeEnv, err) + } + + return func() error { return os.RemoveAll(dir) }, nil +} + +// saveConnection saves remote and key as shard setup would, under a configuration directory of the test's own. +func saveConnection(t *testing.T, remote, key string) { + t.Helper() + + t.Setenv(client.ConfigHomeEnv, t.TempDir()) + path, err := client.ConfigPath(os.Getenv) + if err != nil { + t.Fatalf("ConfigPath: %v", err) + } + if err := client.SaveConfig(path, client.Config{Remote: remote, APIKey: key}); err != nil { + t.Fatalf("SaveConfig: %v", err) + } +} + +// A saved connection names the server with no flag and no SHARD_REMOTE or SHARD_API_KEY, and the verb goes there, not to the socket. (SHARD-657) +func TestASavedConnectionReachesTheFront(t *testing.T) { + var out bytes.Buffer + + app, f, _ := newFrontApp(t, &out) + noRemoteEnv(t) + t.Setenv(client.CAFileEnv, f.ca) + saveConnection(t, f.url, f.key) + // The front still serves the daemon of the old root; the CLI gets a root with no socket at all. + app.Root = t.TempDir() + + if err := app.Run(t.Context(), []string{"list"}); err != nil { + t.Fatalf("list through the saved connection: %v", err) + } + if !strings.Contains(out.String(), "up-1") { + t.Errorf("list through the saved connection printed %q, want the sandbox the daemon holds", out.String()) + } +} + +// An explicit empty --remote asks for the socket, so the saved server sees no dial and the saved file is not read. (SHARD-657) +func TestAnEmptyRemoteFlagIgnoresTheSavedConnection(t *testing.T) { + var out bytes.Buffer + accepted := acceptCount(t) + + app := newListApp(t, &out, listed(), nil) + noRemoteEnv(t) + saveConnection(t, "https://"+accepted.address, leakKey) + + if err := app.Run(t.Context(), []string{"--remote", "", "list"}); err != nil { + t.Fatalf("list over the socket: %v", err) + } + if !strings.Contains(out.String(), "up-1") { + t.Errorf("list over the socket printed %q, want the sandbox the daemon holds", out.String()) + } + accepted.none(t) +} + +// A host verb would report this host as the saved server, so it refuses before it dials or touches the root, and never prints the key. (SHARD-657) +func TestAHostVerbRefusesASavedConnection(t *testing.T) { + accepted := acceptCount(t) + + noRemoteEnv(t) + saveConnection(t, "https://"+accepted.address, leakKey) + for verb, args := range map[string][]string{ + "pull": {"pull", "alpine:3.20"}, + "image list": {"image", "list"}, + "daemon status": {"daemon", "status"}, + "info": {"info"}, + "tokens mint": {"tokens", "mint", "--name", "ci"}, + "tokens list": {"tokens", "list"}, + } { + root := t.TempDir() + app := App{Version: "test", Root: root, Out: io.Discard} + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + err := app.Run(ctx, args) + cancel() + if want := "shard " + verb + " runs on the daemon host only"; err == nil || !strings.Contains(err.Error(), want) || !strings.Contains(err.Error(), "saved connection") { + t.Errorf("%s with a saved connection returned %v, want %q and the saved connection", verb, err, want) + } + if err != nil && strings.Contains(err.Error(), leakKey) { + t.Errorf("%s printed the saved key", verb) + } + entries, err := os.ReadDir(root) + if err != nil { + t.Fatalf("read the root: %v", err) + } + if len(entries) != 0 { + t.Errorf("%s with a saved connection left %d entries in the root, want none", verb, len(entries)) + } + } + + accepted.none(t) +} + +// The saved connection says where commands go, never where a front runs, so serve fails for its own reason alone. (SHARD-657) +func TestServeIgnoresTheSavedConnection(t *testing.T) { + accepted := acceptCount(t) + + noRemoteEnv(t) + saveConnection(t, "https://"+accepted.address, leakKey) + missing := filepath.Join(t.TempDir(), "missing-signing-key") + app := App{Version: "test", Root: t.TempDir(), Out: io.Discard} + + err := app.Run(t.Context(), []string{"serve", "--listen", "127.0.0.1:0", "--signing-key-file", missing}) + if err == nil || strings.Contains(err.Error(), "daemon host only") || !strings.Contains(err.Error(), missing) { + t.Errorf("serve with a saved connection returned %v, want the missing signing key", err) + } + accepted.none(t) +} + +// A saved file no client can use stops every verb with its path, rather than a quiet fall back to the socket. (SHARD-657) +func TestABrokenSavedConnectionNamesTheFile(t *testing.T) { + var out bytes.Buffer + + app := newListApp(t, &out, listed(), nil) + noRemoteEnv(t) + t.Setenv(client.ConfigHomeEnv, t.TempDir()) + path, err := client.ConfigPath(os.Getenv) + if err != nil { + t.Fatalf("ConfigPath: %v", err) + } + if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil { + t.Fatalf("make the configuration directory: %v", err) + } + if err := os.WriteFile(path, []byte(`{"remote":"ftp://shard.example.com","api_key":"`+leakKey+`"}`), 0o600); err != nil { + t.Fatalf("write the configuration: %v", err) + } + + err = app.Run(t.Context(), []string{"list"}) + if err == nil || !strings.Contains(err.Error(), path) || strings.Contains(err.Error(), leakKey) { + t.Errorf("list with a broken saved connection returned %v, want its path and never the key", err) + } + if strings.Contains(out.String(), "up-1") { + t.Error("list with a broken saved connection read the socket") + } +} + +// bearerFront is a plain http server that only records the bearer of each request, so a row can tell which server a verb dialed. +func bearerFront(t *testing.T) (string, func() []string) { + t.Helper() + + var mu sync.Mutex + var bearers []string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + mu.Lock() + bearers = append(bearers, strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ")) + mu.Unlock() + http.Error(w, `{"error":"recorded"}`, http.StatusUnauthorized) + })) + t.Cleanup(server.Close) + + return server.URL, func() []string { + mu.Lock() + defer mu.Unlock() + + return slices.Clone(bearers) + } +} + +// The §13 order: --remote, then SHARD_REMOTE, then the saved connection, then the socket; an explicit empty --remote is the socket. (SHARD-657) +func TestTheRemoteComesFromTheFlagThenTheEnvironmentThenTheSavedConnection(t *testing.T) { + const envKey, savedKey = "shard657-synthetic-env-key", "shard657-synthetic-saved-key" + for _, tc := range []struct { + name string + flag, env, saved, key bool + emptyFlag bool + wantServer, wantBearer string + }{ + {name: "flag over env and saved", flag: true, env: true, saved: true, key: true, wantServer: "flag", wantBearer: envKey}, + {name: "flag over saved", flag: true, saved: true, key: true, wantServer: "flag", wantBearer: envKey}, + {name: "env over saved", env: true, saved: true, key: true, wantServer: "env", wantBearer: envKey}, + {name: "env alone", env: true, key: true, wantServer: "env", wantBearer: envKey}, + {name: "saved alone, with its key", saved: true, wantServer: "saved", wantBearer: savedKey}, + {name: "SHARD_API_KEY over the saved key", saved: true, key: true, wantServer: "saved", wantBearer: envKey}, + {name: "empty flag over env and saved", emptyFlag: true, env: true, saved: true, key: true}, + {name: "none"}, + } { + t.Run(tc.name, func(t *testing.T) { + var out bytes.Buffer + app := newListApp(t, &out, listed(), nil) + noRemoteEnv(t) + t.Setenv(client.ConfigHomeEnv, t.TempDir()) + + urls := map[string]string{} + bearers := map[string]func() []string{} + for _, name := range []string{"flag", "env", "saved"} { + urls[name], bearers[name] = bearerFront(t) + } + if tc.saved { + saveConnection(t, urls["saved"], savedKey) + } + if tc.env { + t.Setenv(client.RemoteEnv, urls["env"]) + } + if tc.key { + t.Setenv(client.APIKeyEnv, envKey) + } + args := []string{"list"} + if tc.flag { + args = append([]string{"--remote", urls["flag"]}, args...) + } + if tc.emptyFlag { + args = append([]string{"--remote", ""}, args...) + } + + err := app.Run(t.Context(), args) + for name, got := range bearers { + if name != tc.wantServer && len(got()) != 0 { + t.Errorf("the %s server got %d requests, want none", name, len(got())) + } + } + if tc.wantServer == "" { + if err != nil || !strings.Contains(out.String(), "up-1") { + t.Errorf("list returned %v and printed %q, want the sandbox the socket holds", err, out.String()) + } + return + } + got := bearers[tc.wantServer]() + if len(got) == 0 || got[0] != tc.wantBearer { + t.Errorf("the %s server got bearers %q, want %q", tc.wantServer, got, tc.wantBearer) + } + if strings.Contains(out.String(), "up-1") { + t.Error("list read the socket, want the remote") + } + }) + } +} diff --git a/cli/daemon.go b/cli/daemon.go index 999bc7d1b..ea8fe8ac2 100644 --- a/cli/daemon.go +++ b/cli/daemon.go @@ -26,7 +26,7 @@ func (a App) daemon(ctx context.Context, args []string) error { if flags.NArg() != 0 { return fmt.Errorf("daemon takes no arguments, or status, got %s", gotArgs(flags.Args())) } - if err := a.hostOnly("daemon"); err != nil { + if err := a.noRemote("daemon"); err != nil { return err } diff --git a/cli/help.go b/cli/help.go index e1a66541e..58af45f35 100644 --- a/cli/help.go +++ b/cli/help.go @@ -26,6 +26,8 @@ type verbHelp struct { summary string // about is the sentence its own help opens with, when it says more than the summary does. about string + // intro is the paragraph after about, ahead of the options. + intro []string args []row flags []flagHelp // env is the variables it reads, which only the top level lists. @@ -33,6 +35,8 @@ type verbHelp struct { notes []note // examples print as written, one call per line, so each pastes whole however wide it is. examples []string + // spaced puts a blank line between the examples. + spaced bool } // row is one line of a two-column list: an argument, a flag or a verb, and what it is. @@ -86,7 +90,7 @@ var verbGroups = []struct { }{ {"Sandboxes", []string{"create", "run", "exec", "list", "logs", "inspect", "stop", "start", "remove", "pause", "resume", "fork", "cp"}}, {"Images, snapshots, secrets and network policies", []string{"pull", "image", "snapshot", "secret", "policy"}}, - {"Host and access", []string{"capabilities", "daemon", "info", "serve", "tokens", "version"}}, + {"Host and access", []string{"capabilities", "daemon", "info", "serve", "setup", "tokens", "version"}}, } // sandboxFlagHelps are the flags create and run share, as sandboxFlags parses them. @@ -126,7 +130,7 @@ var ( }} limitsNote = note{title: "Resource limits", lines: []string{"Use sizes such as 512MiB or 2GiB, and whole numbers for CPUs."}} signingKeyNote = para("The default signing key is created automatically on first use.", "A custom signing key file must already exist.") - hostOnlyNote = para("Runs only on the daemon host and refuses --remote and " + client.RemoteEnv + ".") + hostOnlyNote = para("Runs only on the daemon host and refuses --remote, " + client.RemoteEnv + " and a saved connection.") ) const noPolicyLine = "Without a policy, the sandbox can access the internet but not private networks." @@ -147,7 +151,9 @@ var helps = map[string]verbHelp{ {client.CAFileEnv, "custom CA certificate file; HTTPS only"}, }, notes: []note{ - para(fmt.Sprintf("Set %s and %s for remote access.", client.RemoteEnv, client.APIKeyEnv), "Without a remote URL, Shard connects to the local daemon."), + para(fmt.Sprintf("Set %s and %s for remote access.", client.RemoteEnv, client.APIKeyEnv), "Without a remote URL or a saved connection, Shard connects to the local daemon."), + {title: "Get started", lines: []string{"shard setup"}}, + {title: "For automated setup options", lines: []string{"shard setup --help"}}, }, }, "create": { @@ -568,6 +574,27 @@ var helps = map[string]verbHelp{ notes: []note{para("Shows the version, provider, process details and background tasks."), hostOnlyNote}, examples: []string{"shard daemon status", "shard daemon status --format json"}, }, + "setup": { + usage: []string{"setup [OPTIONS]"}, + summary: "configure a local sandbox host or a remote connection", + about: "Set up Shard on this machine or connect to a remote server.", + intro: []string{"Run without options to start the interactive wizard.", "For automated setup, specify the choices below."}, + flags: []flagHelp{ + {"--local", "Set up this machine to run sandboxes", ""}, + {"--remote ", "Connect to a remote Shard server", ""}, + {"--provider ", "firecracker, gvisor, sysbox, runc or vz", ""}, + {"--start-at-boot ", "Configure automatic daemon startup", ""}, + {"--save", "Save the remote connection", ""}, + {"-y, --yes", "Apply changes without confirmation", ""}, + }, + notes: []note{{title: "Remote authentication", lines: []string{"Set " + client.APIKeyEnv + ". Do not pass the key as a command argument."}}}, + examples: []string{ + "shard setup", + "shard setup --local --provider gvisor --start-at-boot=true -y", + "shard setup --remote https://shard.example.com --save -y", + }, + spaced: true, + }, "capabilities": { usage: []string{"capabilities [OPTIONS]"}, summary: "show the lifecycle verbs the server supports", @@ -689,6 +716,9 @@ func helpText(key string) string { usage = append(usage, lead+line) } sections := []string{strings.Join(usage, "\n"), wrap("", 0, h.opening())} + if len(h.intro) > 0 { + sections = append(sections, wrapLines("", 0, h.intro)) + } if key == "" { sections = append(sections, topLevel()...) @@ -720,7 +750,11 @@ func helpText(key string) string { if len(h.examples) == 1 { heading = "Example:\n " } - sections = append(sections, heading+strings.Join(h.examples, "\n ")) + between := "\n " + if h.spaced { + between = "\n\n " + } + sections = append(sections, heading+strings.Join(h.examples, between)) } return strings.Join(sections, "\n\n") diff --git a/cli/help_test.go b/cli/help_test.go index 53171ca93..784439bcc 100644 --- a/cli/help_test.go +++ b/cli/help_test.go @@ -113,7 +113,7 @@ func TestAVerbUnderABadRemoteFailsOnTheMissingKey(t *testing.T) { } // The top level ends with the global options, the variables a remote reads and how to reach one, as SHARD-503 words them. -func TestTheTopLevelEndsWithTheRemoteSetup(t *testing.T) { +func TestTheTopLevelEndsWithTheRemoteSetupAndGetStarted(t *testing.T) { want := `Global options: --root directory for local Shard data (default ` + DefaultRoot + `) --remote URL of the Shard API server; HTTP/HTTPS supported, @@ -126,7 +126,13 @@ Environment variables: SHARD_CA_FILE custom CA certificate file; HTTPS only Set SHARD_REMOTE and SHARD_API_KEY for remote access. -Without a remote URL, Shard connects to the local daemon. +Without a remote URL or a saved connection, Shard connects to the local daemon. + +Get started: + shard setup + +For automated setup options: + shard setup --help Run 'shard COMMAND --help' for options and examples.` diff --git a/cli/main_integration_test.go b/cli/main_integration_test.go index ffb6b28ac..f5498f5b8 100644 --- a/cli/main_integration_test.go +++ b/cli/main_integration_test.go @@ -49,7 +49,7 @@ var ( ) func TestMain(m *testing.M) { - code, err := run(m) + code, err := isolated(m) if err != nil { fmt.Fprintf(os.Stderr, "cli integration tests: %v\n", err) } @@ -57,6 +57,21 @@ func TestMain(m *testing.M) { os.Exit(code) } +// isolated is run under a saved connection of its own, which the shard binary under test inherits, so no run reads the user's. +func isolated(m *testing.M) (int, error) { + cleanup, err := isolateConfig() + if err != nil { + return 1, err + } + + code, err := run(m) + if cleanupErr := cleanup(); cleanupErr != nil { + return 1, errors.Join(err, cleanupErr) + } + + return code, err +} + // run gives the package one daemon of its own, so no test speaks to the daemon of the systemd unit. func run(m *testing.M) (int, error) { if !hostRunsSandboxes() { diff --git a/cli/main_test.go b/cli/main_test.go new file mode 100644 index 000000000..4f9968139 --- /dev/null +++ b/cli/main_test.go @@ -0,0 +1,25 @@ +//go:build !integration + +package cli + +import ( + "fmt" + "os" + "testing" +) + +func TestMain(m *testing.M) { + cleanup, err := isolateConfig() + if err != nil { + fmt.Fprintf(os.Stderr, "cli tests: %v\n", err) + os.Exit(1) + } + + code := m.Run() + if err := cleanup(); err != nil { + fmt.Fprintf(os.Stderr, "cli tests: %v\n", err) + code = 1 + } + + os.Exit(code) +} diff --git a/cli/serve.go b/cli/serve.go index f9f35c0b8..be07082b9 100644 --- a/cli/serve.go +++ b/cli/serve.go @@ -19,7 +19,7 @@ func (a App) serve(ctx context.Context, args []string) error { if flags.NArg() != 0 { return fmt.Errorf("serve takes no arguments, got %s", gotArgs(flags.Args())) } - if err := a.hostOnly("serve"); err != nil { + if err := a.noRemote("serve"); err != nil { return err } diff --git a/cli/setup.go b/cli/setup.go new file mode 100644 index 000000000..125b90376 --- /dev/null +++ b/cli/setup.go @@ -0,0 +1,244 @@ +package cli + +import ( + "cmp" + "context" + "errors" + "fmt" + "os" + "slices" + + "github.com/presmihaylov/shard/pkg/term" + "github.com/presmihaylov/shard/services/client" + "github.com/presmihaylov/shard/services/setup" +) + +// setupFlags are the choices a command line gives shard setup ahead of the wizard. +type setupFlags struct { + local bool + remote string + provider string + startAtBoot boolChoice + save bool + yes bool +} + +// boolChoice is a flag that takes true or false as its value, so a bare --start-at-boot is refused. +type boolChoice struct { + set bool + value bool +} + +func (b *boolChoice) String() string { + if !b.set { + return "" + } + + return fmt.Sprint(b.value) +} + +func (b *boolChoice) Set(value string) error { + switch value { + case "true": + b.set, b.value = true, true + case "false": + b.set, b.value = true, false + default: + return fmt.Errorf("want true or false, got %q", value) + } + + return nil +} + +func (a App) setup(ctx context.Context, args []string) error { + var opts setupFlags + flags := newFlags("setup") + flags.BoolVar(&opts.local, "local", false, "") + flags.StringVar(&opts.remote, "remote", "", "") + flags.StringVar(&opts.provider, "provider", "", "") + flags.Var(&opts.startAtBoot, "start-at-boot", "") + flags.BoolVar(&opts.save, "save", false, "") + flags.BoolVar(&opts.yes, "y", false, "") + flags.BoolVar(&opts.yes, "yes", false, "") + if err := parseVerb(flags, args); err != nil { + return err + } + if flags.NArg() > 0 { + return fmt.Errorf("setup takes no arguments, got %s", gotArgs(flags.Args())) + } + if err := opts.check(); err != nil { + return err + } + + host, err := setup.NewHost(a.Version) + if err != nil { + return err + } + run := setup.Setup{Host: host, UI: &answers{opts: opts, t: term.New(a.stdin(), a.Out, os.Getenv), env: os.Getenv}} + + return setupExit(run.Run(ctx)) +} + +// setupExit is the exit a setup error takes: a stopped step already said why under its checklist, and an interrupt leaves as SIGINT does. +func setupExit(err error) error { + interrupted := errors.Is(err, term.ErrInterrupted) || errors.Is(err, context.Canceled) + var stopped *setup.StoppedError + switch { + case err == nil: + return nil + case errors.As(err, &stopped) && interrupted: + return &ExitError{Code: InterruptedExitCode} + case errors.As(err, &stopped): + return &ExitError{Code: 1} + case interrupted: + return &ExitError{Code: InterruptedExitCode, Message: "setup interrupted"} + } + + return err +} + +// check refuses a local choice beside a remote one, so neither half guesses which was meant. +func (o setupFlags) check() error { + localOnly := o.provider != "" || o.startAtBoot.set + if o.remote != "" && (o.local || localOnly) { + return errors.New("--remote connects to a server; --local, --provider and --start-at-boot set up this machine") + } + if o.save && (o.local || localOnly) { + return errors.New("--save saves a remote connection; it does not apply to --local") + } + + return nil +} + +// answers puts the flags in front of the terminal: a flag answers its question, else a person does, else the error names the flag. +type answers struct { + opts setupFlags + t *term.Terminal + env func(string) string + // urlAsked is set once the URL was answered, so an edit after a failed check asks the person and never loops on the flag. + urlAsked bool +} + +// flagged is the option name a flag picks for a question, and whether one does. +func (a *answers) flagged(q setup.Question) (string, bool) { + switch q { + case setup.AskMode: + if a.opts.remote != "" || a.opts.save { + return "remote", true + } + if a.opts.local || a.opts.provider != "" || a.opts.startAtBoot.set { + return "local", true + } + case setup.AskProvider: + return a.opts.provider, a.opts.provider != "" + case setup.AskStartAtBoot: + return a.opts.startAtBoot.String(), a.opts.startAtBoot.set + case setup.AskSaved: + return "replace", a.opts.remote != "" + case setup.AskRetry: + // A failed check with nobody to ask leaves with its reason rather than an error about the terminal. + return "exit", !a.t.Interactive() + } + + return "", false +} + +func (a *answers) Select(ctx context.Context, q setup.Question, title string, options []term.Option) (int, error) { + if name, ok := a.flagged(q); ok { + return pick(q, options, name) + } + chosen, err := a.t.Select(ctx, title, options) + + return chosen, need(q, err) +} + +// pick is the option a flag names; an unavailable one is refused with its reason, never swapped for another. +func pick(q setup.Question, options []term.Option, name string) (int, error) { + var names []string + for i, o := range options { + if o.Name != name { + names = append(names, o.Name) + continue + } + if len(o.Unavailable) > 0 { + return 0, fmt.Errorf("%s %s: %s is unavailable: %s", setupFlag[q], name, o.Label, o.Unavailable[0]) + } + + return i, nil + } + + return 0, fmt.Errorf("%s %q: want %s", setupFlag[q], name, orList(names)) +} + +func (a *answers) Confirm(ctx context.Context, q setup.Question, text string, yes bool) (bool, error) { + switch { + case slices.Contains(confirmations, q) && a.opts.yes: + return true, nil + case q == setup.AskSave && (a.opts.save || a.opts.yes || !a.t.Interactive()): + return a.opts.save, nil + case !slices.Contains(confirmations, q) && !a.t.Interactive(): + // A choice no flag names keeps things as they are when nobody is there to ask. + return false, nil + } + answer, err := a.t.Confirm(ctx, text, yes) + + return answer, need(q, err) +} + +func (a *answers) Text(ctx context.Context, q setup.Question, prompt string) (string, error) { + if q == setup.AskURL && !a.urlAsked { + a.urlAsked = true + if url := cmp.Or(a.opts.remote, a.env(client.RemoteEnv)); url != "" { + return url, nil + } + } + answer, err := a.t.Text(ctx, prompt) + + return answer, need(q, err) +} + +func (a *answers) Secret(ctx context.Context, q setup.Question, prompt string) (string, error) { + answer, err := a.t.Secret(ctx, prompt) + + return answer, need(q, err) +} + +func (a *answers) Checklist(title string, steps []string) (setup.Checklist, error) { + list, err := a.t.Checklist(title, steps) + if err != nil { + return nil, err + } + + return list, nil +} + +func (a *answers) Print(lines ...string) error { return a.t.Print(lines...) } + +// confirmations are the questions -y answers; the rest are choices it never makes. +var confirmations = []setup.Question{setup.AskConfirm, setup.AskHTTP} + +// setupFlag is the option that answers each question, which a run without a terminal must name. +var setupFlag = map[setup.Question]string{ + setup.AskMode: "--local or --remote ", + setup.AskProvider: "--provider", + setup.AskStartAtBoot: "--start-at-boot", + setup.AskConfirm: "-y", + setup.AskHTTP: "-y", + setup.AskURL: "--remote", + setup.AskAPIKey: client.APIKeyEnv, +} + +// need words a question asked without a terminal as the option that answers it. +func need(q setup.Question, err error) error { + if !errors.Is(err, term.ErrNotTerminal) { + return err + } + switch flag, ok := setupFlag[q]; { + case q == setup.AskAPIKey: + return fmt.Errorf("no terminal to read the API key: set %s", flag) + case ok: + return fmt.Errorf("no terminal to ask %s: pass %s", q, flag) + } + + return fmt.Errorf("no terminal to ask %s: run shard setup in a terminal", q) +} diff --git a/cli/setup_test.go b/cli/setup_test.go new file mode 100644 index 000000000..a0698d28f --- /dev/null +++ b/cli/setup_test.go @@ -0,0 +1,161 @@ +package cli + +import ( + "bytes" + "context" + "errors" + "fmt" + "os" + "strings" + "testing" + + "github.com/presmihaylov/shard/pkg/term" + "github.com/presmihaylov/shard/services/client" + "github.com/presmihaylov/shard/services/setup" +) + +func TestSetupRefusesWhatItCannotRun(t *testing.T) { + for _, tc := range []struct { + args []string + want string + }{ + {[]string{"setup", "now"}, "setup takes no arguments"}, + {[]string{"setup", "--remote", "https://shard.example.com", "--provider", "gvisor"}, "--remote connects to a server"}, + {[]string{"setup", "--remote", "https://shard.example.com", "--local"}, "--remote connects to a server"}, + {[]string{"setup", "--save", "--start-at-boot=false"}, "--save saves a remote connection"}, + {[]string{"setup", "--start-at-boot=yes"}, "want true or false"}, + } { + err := (App{Version: "test", Root: t.TempDir(), Out: &bytes.Buffer{}}).run(t.Context(), tc.args) + if err == nil || !strings.Contains(err.Error(), tc.want) { + t.Errorf("shard %s: %v, want %q", strings.Join(tc.args, " "), err, tc.want) + } + } +} + +// noTerminal is the answers of these flags with stdin a pipe, the way a script runs setup. +func noTerminal(t *testing.T, opts setupFlags) *answers { + t.Helper() + in, w, err := os.Pipe() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + if err := errors.Join(in.Close(), w.Close()); err != nil { + t.Error(err) + } + }) + + none := func(string) string { return "" } + + return &answers{opts: opts, t: term.New(in, &bytes.Buffer{}, none), env: none} +} + +var providers = []term.Option{ + {Name: "firecracker", Label: "Firecracker", Unavailable: []string{"/dev/kvm is missing"}}, + {Name: "gvisor", Label: "gVisor"}, +} + +func TestAFlagPicksItsOptionOrSaysWhyNot(t *testing.T) { + for _, tc := range []struct { + provider string + want int + err string + }{ + {"gvisor", 1, ""}, + {"firecracker", 0, "--provider firecracker: Firecracker is unavailable: /dev/kvm is missing"}, + {"kata", 0, `--provider "kata": want firecracker or gvisor`}, + } { + chosen, err := noTerminal(t, setupFlags{provider: tc.provider}).Select(t.Context(), setup.AskProvider, "Provider", providers) + if tc.err == "" && (err != nil || chosen != tc.want) { + t.Errorf("--provider %s chose %d, %v", tc.provider, chosen, err) + } + if tc.err != "" && (err == nil || err.Error() != tc.err) { + t.Errorf("--provider %s: %v, want %q", tc.provider, err, tc.err) + } + } +} + +func TestWithoutATerminalTheErrorNamesWhatAnswers(t *testing.T) { + ui := noTerminal(t, setupFlags{}) + _, provider := ui.Select(t.Context(), setup.AskProvider, "Provider", providers) + _, key := ui.Secret(t.Context(), setup.AskAPIKey, "API key") + _, existing := ui.Select(t.Context(), setup.AskExisting, "What would you like to do?", providers) + + for got, want := range map[error]string{ + provider: "no terminal to ask provider: pass --provider", + key: "no terminal to read the API key: set SHARD_API_KEY", + existing: "no terminal to ask existing: run shard setup in a terminal", + } { + if got == nil || got.Error() != want { + t.Errorf("got %v, want %q", got, want) + } + } +} + +func TestYesConfirmsButMakesNoChoice(t *testing.T) { + ui := noTerminal(t, setupFlags{local: true, yes: true}) + for q, want := range map[setup.Question]bool{setup.AskConfirm: true, setup.AskHTTP: true, setup.AskSave: false, setup.AskSwitch: false} { + got, err := ui.Confirm(t.Context(), q, "?", true) + if err != nil || got != want { + t.Errorf("-y answered %s with %v, %v; want %v", q, got, err, want) + } + } + if mode, ok := ui.flagged(setup.AskMode); !ok || mode != "local" { + t.Errorf("--local picked %q", mode) + } +} + +func TestTheURLIsTheFlagOrTheEnvironmentOnlyOnce(t *testing.T) { + for _, tc := range []struct { + flag, env, want string + }{ + {"https://flag.example.com", "https://env.example.com", "https://flag.example.com"}, + {"", "https://env.example.com", "https://env.example.com"}, + } { + ui := noTerminal(t, setupFlags{remote: tc.flag}) + ui.env = func(name string) string { return map[string]string{client.RemoteEnv: tc.env}[name] } + if url, err := ui.Text(t.Context(), setup.AskURL, "Shard server URL:"); err != nil || url != tc.want { + t.Errorf("the URL is %q, %v; want %q", url, err, tc.want) + } + if _, err := ui.Text(t.Context(), setup.AskURL, "Shard server URL:"); err == nil || err.Error() != "no terminal to ask url: pass --remote" { + t.Errorf("an edit after a failed check answered %v, want the person asked", err) + } + } +} + +func TestARemoteFlagReplacesASavedConnectionAndAFailedCheckExits(t *testing.T) { + ui := noTerminal(t, setupFlags{remote: "https://shard.example.com", save: true, yes: true}) + saved := []term.Option{{Name: "check"}, {Name: "replace"}, {Name: "remove"}, {Name: "exit"}} + if chosen, err := ui.Select(t.Context(), setup.AskSaved, "What would you like to do?", saved); err != nil || saved[chosen].Name != "replace" { + t.Errorf("--remote chose %d, %v; want replace", chosen, err) + } + retry := []term.Option{{Name: "retry"}, {Name: "edit"}, {Name: "exit"}} + if chosen, err := ui.Select(t.Context(), setup.AskRetry, "What would you like to do?", retry); err != nil || retry[chosen].Name != "exit" { + t.Errorf("a failed check without a terminal chose %d, %v; want exit", chosen, err) + } +} + +func TestSetupExitsOnceItSaidWhy(t *testing.T) { + refused := errors.New("--provider kata: want firecracker or gvisor") + for _, tc := range []struct { + err error + code int + message string + }{ + {&setup.StoppedError{Step: "Install gVisor", Err: errors.New("timed out")}, 1, ""}, + {&setup.StoppedError{Step: "Install gVisor", Err: context.Canceled}, InterruptedExitCode, ""}, + {fmt.Errorf("select mode: %w", term.ErrInterrupted), InterruptedExitCode, "setup interrupted"}, + {fmt.Errorf("read the terminal: %w", context.Canceled), InterruptedExitCode, "setup interrupted"}, + } { + var exit *ExitError + if !errors.As(setupExit(tc.err), &exit) || exit.Code != tc.code || exit.Message != tc.message { + t.Errorf("%v exits %+v, want code %d and %q", tc.err, exit, tc.code, tc.message) + } + } + if err := setupExit(refused); !errors.Is(err, refused) { + t.Errorf("a refusal became %v", err) + } + if err := setupExit(nil); err != nil { + t.Errorf("a clean run exits with %v", err) + } +} diff --git a/docs/cli.md b/docs/cli.md index 3c2ccc089..bcf5ce890 100644 --- a/docs/cli.md +++ b/docs/cli.md @@ -14,6 +14,7 @@ This is the final shape of every verb, flag and output of `shard`. The SDKs buil | 128 + n | `run` of an app that signal n ended | | 125 | `run` when shard itself fails, so a script tells it from an app that exits 1 | | 130 | `run` that Ctrl+C left; the sandbox stays running | +| 130 | `setup` that Ctrl+C left; the steps done so far stay in place | `daemon status` exits 1 when a background task is in backoff, after it prints the whole status. @@ -33,8 +34,8 @@ A remote client reads three environment variables. | variable | what | | --- | --- | -| `SHARD_REMOTE` | the API server URL; `--remote` overrides it | -| `SHARD_API_KEY` | the API token, the `token` field of a `shard tokens mint` record | +| `SHARD_REMOTE` | the API server URL; `--remote` overrides it, and it overrides the connection `shard setup` saved | +| `SHARD_API_KEY` | the API token, the `token` field of a `shard tokens mint` record; it overrides the saved key | | `SHARD_CA_FILE` | a custom CA certificate file, for `https` only; unset, the host's trust store decides | An `http` remote encrypts nothing, so every command over one prints one warning to stderr and never @@ -44,8 +45,8 @@ to stdout. Use it only on localhost or through a trusted encrypted network. `SHA `pull`, `image list`, `image remove`, `image prune` and `daemon status` run on the daemon host only, because `shard serve` refuses their routes. `daemon`, `serve`, `info`, `tokens mint`, `tokens list` and `tokens revoke` act on the files and the processes of this host, so they run there only too. -With `--remote` or `SHARD_REMOTE` set, each one fails before it dials or touches the host, and its -error names the verb. `tokens scopes` asks the server, so it follows the remote. +With `--remote`, `SHARD_REMOTE` or a saved connection, each one fails before it dials or touches the +host, and its error names the verb. `--remote ""` runs one on this host past a saved connection. `tokens scopes` asks the server, so it follows the remote. ## Names and aliases @@ -125,6 +126,7 @@ shard: sandbox is paused: resume it with shard resume | `daemon status` | `--format` | table | the daemon's state and its tasks | | `info` | `--format` | table | the provider a daemon would pick, and why | | `serve` | `--listen --signing-key-file` | - | its log | +| `setup` | `--local --remote --provider --start-at-boot --save -y/--yes` | - | the wizard, or the checklist of a run with every answer given; see [setup](setup.md) | | `tokens mint` | `--name --signing-key-file --duration --scopes --format` | json | the token record | | `tokens list` | `--signing-key-file --format` | table | the ledger | | `tokens revoke ` | `--name --signing-key-file` | - | `revoked token `, or `revoked tokens of ` with `--name` | diff --git a/docs/setup.md b/docs/setup.md new file mode 100644 index 000000000..7dc9ad1d1 --- /dev/null +++ b/docs/setup.md @@ -0,0 +1,152 @@ +# Set up Shard + +`shard setup` prepares this machine to run sandboxes, or connects the CLI to a remote Shard server. +Run it with no options and it asks each question in turn. Every question has an option, so a script +takes the same path a person does, through the same checks. + +```sh +shard setup +``` + +## Supported platforms + +| host | providers | background service | +| --- | --- | --- | +| Linux x86-64 | `firecracker` (needs `/dev/kvm`), `gvisor`, `sysbox`, `runc` | systemd | +| macOS 14 or later on Apple silicon | `vz` | launchd | + +The wizard always lists all five providers. One this host cannot run is shown as unavailable, with +the reason, and cannot be chosen. A runtime setup can install does not make a provider unavailable. +Any other host, such as an Intel Mac, can still connect to a remote server. + +## Local setup + +The wizard asks for a provider and whether the daemon starts at boot, then checks the host, shows +what it will change, and asks once before it changes anything. + +- **Provider.** On Linux it recommends Firecracker when `/dev/kvm` is usable and gVisor otherwise. + On a supported Mac it recommends macOS Virtualization. A provider that fails setup is never + swapped for another. +- **Automatic startup.** Yes installs and starts a systemd service on Linux or a launchd service on + macOS, which starts the daemon at boot and restarts it after a crash. No installs the provider's + tools and prints the `shard daemon` command to run yourself. +- **Data.** The daemon keeps its state in `/var/lib/shard`. Setup never asks for another directory. +- **Checks.** Setup checks the operating system, the provider's requirements, administrator access, + the install paths, disk space and filesystem, download access, an earlier installation and the + service manager before it changes the host. A failed check changes nothing. +- **Network.** Setup never exposes the HTTP API. `shard serve` stays a separate step; see + [the daemon guide](daemon.md). + +Setup runs as you and uses `sudo` for the steps that need root. It installs the running binary as +`/usr/local/bin/shard`, owned by root, and `shard-init` from the same release, after it checks that +download against the release's `SHA256SUMS`. Setup creates no `shard` group, so on Linux the API +socket belongs to root and local commands run with `sudo`. [The API socket](daemon.md#the-api-socket) +says what a group you create yourself changes. + +Declining at the confirmation leaves the host as it was. + +### Progress and failures + +A live checklist shows each step: a gray circle for a pending step, a blue spinner for the one that +runs, a green check for a step done, yellow for one that needs attention and a red cross for a +failure. Each mark is a symbol as well as a color, and `NO_COLOR` turns the color off. When the +output is not a terminal, setup prints each step once, as it ends. + +On a failure setup marks the step, says why, and stops: + +```text +✗ Install gVisor + Could not download runsc: connection timed out. + +Setup stopped. Earlier completed steps remain in place. +Run `shard setup` again to retry. +``` + +Every step is safe to run again, so a second `shard setup` picks up where the first stopped. + +## An existing installation + +On a host it set up before, `shard setup` shows the version, the provider and the service, and +offers to check or repair the installation, upgrade Shard, uninstall it, or exit. + +- **Check or repair** inspects the tools, the permissions and the service, and shows each change + before it makes one. It keeps the settings and the data. +- **Upgrade** downloads and verifies the new release before it replaces anything, keeps the old + binary until the new one checks out, and says before it restarts the daemon. +- **Uninstall** stops and removes the service and the files setup installed. It refuses while + sandboxes exist and names the commands that remove them. The data in `/var/lib/shard` and any + tool another program may share stay. + +Setup records what it created in `/var/lib/shard-setup/manifest.json`, and removes only what that +file lists. An installation setup did not make is reported as a manual installation; setup then +explains what it found and changes none of it. + +## Remote setup + +The wizard asks for the server URL, the URL of `shard serve` or the proxy in front of it. `https` +is recommended. An `http` URL works only after a warning that it encrypts neither the API key nor +the requests. + +It reads the API key from `SHARD_API_KEY` when that is set, and otherwise asks for it with hidden +input. The key never appears on the screen, in a log or in an error. Then it checks that it can +reach the server, that the key is accepted, and which lifecycle verbs the server supports. A failed +check offers to retry, to edit the details, or to exit. + +Certificates are always verified. For a server whose certificate a private CA signed, set +`SHARD_CA_FILE` to that CA's certificate. There is no option to skip verification. + +### The saved connection + +After a successful check, setup offers to save the connection so every later command uses it. The +file is `$XDG_CONFIG_HOME/shard/config.json`, or `~/.config/shard/config.json` when that is unset: + +```json +{ + "remote": "https://shard.example.com", + "api_key": "" +} +``` + +The key is stored as plain text, so setup makes the directory and the file readable by your user +only, and replaces the file whole. A command resolves its server and key in this order: + +| setting | first | then | last | +| --- | --- | --- | --- | +| server URL | `--remote` | `SHARD_REMOTE` | the saved connection | +| API key | `SHARD_API_KEY` | | the saved connection | + +A verb that runs on the daemon host only refuses a saved connection, and `--remote ""` runs it on +this host. Run `shard setup` again to check, replace or remove the saved connection. Choosing local +setup while a connection is saved offers to remove it, and removes it only once local setup +succeeds; a repair or an upgrade of an existing installation makes the same offer. A `SHARD_REMOTE` +in the environment still overrides the local daemon after that. + +## Automated setup + +| option | what | +| --- | --- | +| `--local` | set up this machine to run sandboxes | +| `--remote ` | connect to a remote Shard server | +| `--provider ` | `firecracker`, `gvisor`, `sysbox`, `runc` or `vz` | +| `--start-at-boot ` | install the background service, or leave the daemon to you | +| `--save` | save the remote connection | +| `-y`, `--yes` | apply the changes without the confirmation | + +```sh +shard setup --local --provider gvisor --start-at-boot=true -y +SHARD_API_KEY=... shard setup --remote https://shard.example.com --save -y +``` + +An option answers its question and the wizard asks the rest. Without a terminal, every question +must have its answer: setup fails and names the option it needs, and it never picks one for you. +`-y` answers the confirmation and the `http` warning only. It skips no check, and it saves a +connection only with `--save`. The API key comes from `SHARD_API_KEY`, never from an option. Setup +refuses a local option beside `--remote`, and `--save` beside a local option. + +## Exit codes + +| code | when | +| --- | --- | +| 0 | setup finished, or there was nothing to do | +| 1 | a check or a step failed, an option was refused, or the confirmation was declined | +| 130 | Ctrl+C left setup; the steps done so far stay in place | diff --git a/go.mod b/go.mod index 25033a80e..e314d3b2f 100644 --- a/go.mod +++ b/go.mod @@ -13,6 +13,7 @@ require ( github.com/danielgtaylor/huma/v2 v2.39.1 github.com/golang-jwt/jwt/v5 v5.3.1 github.com/google/go-containerregistry v0.21.9 + github.com/klauspost/compress v1.19.1 github.com/moby/profiles/apparmor v0.2.3 github.com/moby/profiles/seccomp v0.2.4 github.com/opencontainers/runtime-spec v1.3.0 @@ -36,7 +37,6 @@ require ( github.com/gogo/protobuf v1.3.2 // indirect github.com/golang/groupcache v0.0.0-20241129210726-2c02b8208cf8 // indirect github.com/google/btree v1.1.2 // indirect - github.com/klauspost/compress v1.19.1 // indirect github.com/moby/sys/sequential v0.6.0 // indirect github.com/moby/sys/userns v0.1.0 // indirect github.com/opencontainers/go-digest v1.0.0 // indirect diff --git a/pkg/firecracker/version.go b/pkg/firecracker/version.go index 0680f3b73..a15d4360a 100644 --- a/pkg/firecracker/version.go +++ b/pkg/firecracker/version.go @@ -32,6 +32,12 @@ func CheckVersion(binary string) error { if err != nil { return fmt.Errorf("read the version of %s: %w", binary, err) } + + return CheckVersionOutput(binary, out) +} + +// CheckVersionOutput is CheckVersion over what binary --version already printed. +func CheckVersionOutput(binary string, out []byte) error { v, err := parseVersion(string(out)) if err != nil { return fmt.Errorf("read the version of %s: %w", binary, err) diff --git a/pkg/term/term.go b/pkg/term/term.go new file mode 100644 index 000000000..f9adf3260 --- /dev/null +++ b/pkg/term/term.go @@ -0,0 +1,549 @@ +// Package term draws the prompts and the live checklist of an interactive verb, and plain lines when +// the input or the output is not a terminal. It knows key codes and escape sequences, nothing of setup. +package term + +import ( + "context" + "errors" + "fmt" + "io" + "os" + "strings" + "sync" + "time" + "unicode/utf8" + + "github.com/presmihaylov/shard/pkg/pty" +) + +// ErrNotTerminal is a prompt asked without a terminal, so the caller names the option that answers it. +var ErrNotTerminal = errors.New("this question needs a terminal") + +// ErrInterrupted is Ctrl-C at a prompt in raw mode, where it reaches shard as a key and not a signal. +var ErrInterrupted = errors.New("interrupted") + +// errInputEnded is the end of the input before the answer did. +var errInputEnded = errors.New("the input ended before an answer") + +// The escape sequences the prompts draw with. +const ( + clearLine = "\x1b[2K" + keyUp = "\x1b[A" + keyDown = "\x1b[B" + ctrlC = 3 + ctrlD = 4 + ctrlU = 21 + backspace = 8 + del = 127 +) + +// The colors, each beside a symbol or a word, so a terminal without color loses no meaning. +const ( + green = "32" + blue = "34" + yellow = "33" + red = "31" + gray = "90" + bold = "1" +) + +// Option is one row of a Select. +type Option struct { + // Name is the word that picks the option from a command line, as gvisor for gVisor. + Name string + Label string + // Lines are the description under the label. + Lines []string + // Unavailable is the reason the row cannot be chosen; any line disables it. + Unavailable []string + // Default is the row the cursor starts on. + Default bool +} + +func (o Option) disabled() bool { return len(o.Unavailable) > 0 } + +// Terminal is one person at one terminal, or a script when the input or the output is not one. +type Terminal struct { + out io.Writer + interactive bool + color bool + // input reads the keys; a cancel ends a read that would block. + input func(ctx context.Context) io.Reader + raw func() (pty.Restore, error) +} + +// New is the terminal on in and out; NO_COLOR set to anything drops the color, as no-color.org says. +func New(in *os.File, out io.Writer, getenv func(string) string) *Terminal { + file, isFile := out.(*os.File) + interactive := isFile && pty.IsTerminal(in) && pty.IsTerminal(file) + + return &Terminal{ + out: out, + interactive: interactive, + color: interactive && getenv("NO_COLOR") == "", + input: func(ctx context.Context) io.Reader { return pty.Input(ctx, in) }, + raw: func() (pty.Restore, error) { return pty.MakeRaw(in) }, + } +} + +// Interactive reports whether a person can answer, so a caller without one asks for an option instead. +func (t *Terminal) Interactive() bool { return t.interactive } + +// Print writes each line as it is. +func (t *Terminal) Print(lines ...string) error { + for _, line := range lines { + if _, err := fmt.Fprintln(t.out, line); err != nil { + return fmt.Errorf("write to the terminal: %w", err) + } + } + + return nil +} + +// Select draws the options with a cursor the arrow keys move past the disabled rows, and returns the chosen index. +func (t *Terminal) Select(ctx context.Context, title string, options []Option) (chosen int, err error) { + if !t.interactive { + return 0, ErrNotTerminal + } + at, ok := start(options) + if !ok { + return 0, errors.New("every option is unavailable") + } + + restore, err := t.raw() + if err != nil { + return 0, err + } + defer func() { + if restoreErr := restore(); restoreErr != nil && err == nil { + err = fmt.Errorf("restore the terminal: %w", restoreErr) + } + }() + + drawn := 0 + keys := t.input(ctx) + for { + if drawn, err = t.redraw(drawn, t.selectLines(title, options, at), "\r\n"); err != nil { + return 0, err + } + key, err := readKey(keys) + if err != nil { + return 0, err + } + switch key { + case "\r", "\n": + return at, nil + case string(rune(ctrlC)): + return 0, ErrInterrupted + case keyUp, "k": + at = step(options, at, -1) + case keyDown, "j": + at = step(options, at, 1) + } + } +} + +// start is the Default row when it can be chosen, or else the first that can. +func start(options []Option) (int, bool) { + for i, o := range options { + if o.Default && !o.disabled() { + return i, true + } + } + for i, o := range options { + if !o.disabled() { + return i, true + } + } + + return 0, false +} + +// step moves the cursor to the next row in that direction that can be chosen, and stays at an end. +func step(options []Option, at, by int) int { + for i := at + by; i >= 0 && i < len(options); i += by { + if !options[i].disabled() { + return i + } + } + + return at +} + +func (t *Terminal) selectLines(title string, options []Option, at int) []string { + lines := []string{title, ""} + described := false + for _, o := range options { + described = described || len(o.Lines)+len(o.Unavailable) > 0 + } + for i, o := range options { + if described && i > 0 { + lines = append(lines, "") + } + lines = append(lines, t.optionLines(o, i == at)...) + } + + return lines +} + +func (t *Terminal) optionLines(o Option, current bool) []string { + label := " " + o.Label + if o.disabled() { + label += " [Unavailable]" + } + switch { + case current: + label = t.paint(bold, "❯ "+o.Label) + case o.disabled(): + label = t.paint(gray, label) + } + + lines := []string{label} + for _, line := range append(append([]string(nil), o.Lines...), o.Unavailable...) { + text := " " + line + if o.disabled() { + text = t.paint(gray, text) + } + lines = append(lines, text) + } + + return lines +} + +// Confirm asks a yes or no question; an empty answer takes yes, which the prompt shows as [Y/n] or [y/N]. +func (t *Terminal) Confirm(ctx context.Context, question string, yes bool) (bool, error) { + if !t.interactive { + return false, ErrNotTerminal + } + hint := "[y/N]" + if yes { + hint = "[Y/n]" + } + + keys := t.input(ctx) + for { + if _, err := fmt.Fprintf(t.out, "%s %s ", question, hint); err != nil { + return false, fmt.Errorf("write to the terminal: %w", err) + } + answer, err := readLine(keys) + if err != nil { + return false, err + } + switch strings.ToLower(strings.TrimSpace(answer)) { + case "": + return yes, nil + case "y", "yes": + return true, nil + case "n", "no": + return false, nil + } + } +} + +// Text asks for one line under the prompt, after a > the way the spec draws it. +func (t *Terminal) Text(ctx context.Context, prompt string) (string, error) { + if !t.interactive { + return "", ErrNotTerminal + } + if _, err := fmt.Fprintf(t.out, "%s\n> ", prompt); err != nil { + return "", fmt.Errorf("write to the terminal: %w", err) + } + answer, err := readLine(t.input(ctx)) + if err != nil { + return "", err + } + + return strings.TrimSpace(answer), nil +} + +// Secret asks for a value that echoes as one dot per character, so it lands on no screen and no scrollback. +func (t *Terminal) Secret(ctx context.Context, prompt string) (secret string, err error) { + if !t.interactive { + return "", ErrNotTerminal + } + if _, err := fmt.Fprintf(t.out, "%s\r\n> ", prompt); err != nil { + return "", fmt.Errorf("write to the terminal: %w", err) + } + + restore, err := t.raw() + if err != nil { + return "", err + } + defer func() { + if restoreErr := restore(); restoreErr != nil && err == nil { + err = fmt.Errorf("restore the terminal: %w", restoreErr) + } + }() + + var value []byte + keys := t.input(ctx) + for { + b, err := readByte(keys) + if err != nil { + return "", err + } + echo := "" + switch { + case b == '\r' || b == '\n': + _, err := io.WriteString(t.out, "\r\n") + + return string(value), wrapWrite(err) + case b == ctrlC: + return "", ErrInterrupted + case b == ctrlD && len(value) == 0: + return "", errInputEnded + case b == ctrlU: + echo = strings.Repeat("\b \b", utf8.RuneCount(value)) + value = value[:0] + case b == del || b == backspace: + if len(value) > 0 { + _, size := utf8.DecodeLastRune(value) + value = value[:len(value)-size] + echo = "\b \b" + } + case b >= ' ': + value = append(value, b) + // One dot per character, so the dots count runes and not bytes. + if !utf8.RuneStart(b) { + break + } + echo = "•" + } + if _, err := io.WriteString(t.out, echo); err != nil { + return "", wrapWrite(err) + } + } +} + +func wrapWrite(err error) error { + if err != nil { + return fmt.Errorf("write to the terminal: %w", err) + } + + return nil +} + +// redraw replaces the lines it drew last time, and returns how many it drew now. +func (t *Terminal) redraw(drawn int, lines []string, newline string) (int, error) { + var b strings.Builder + if drawn > 0 { + fmt.Fprintf(&b, "\r\x1b[%dA", drawn) + } + for _, line := range lines { + b.WriteString(clearLine + line + newline) + } + if _, err := io.WriteString(t.out, b.String()); err != nil { + return 0, wrapWrite(err) + } + + return len(lines), nil +} + +func (t *Terminal) paint(color, s string) string { + if !t.color { + return s + } + + return "\x1b[" + color + "m" + s + "\x1b[0m" +} + +// readKey reads one key: a byte, or the three bytes of an arrow. +func readKey(r io.Reader) (string, error) { + b, err := readByte(r) + if err != nil || b != 0x1b { + return string(rune(b)), err + } + seq := []byte{b} + for range 2 { + next, err := readByte(r) + if err != nil { + return "", err + } + seq = append(seq, next) + } + + return string(seq), nil +} + +// readLine reads up to a newline one byte at a time, so nothing past the answer is buffered away from the next prompt. +func readLine(r io.Reader) (string, error) { + var line []byte + for { + b, err := readByte(r) + if errors.Is(err, errInputEnded) && len(line) > 0 { + return string(line), nil + } + if err != nil { + return "", err + } + if b == '\n' { + return strings.TrimSuffix(string(line), "\r"), nil + } + line = append(line, b) + } +} + +func readByte(r io.Reader) (byte, error) { + var b [1]byte + n, err := r.Read(b[:]) + if n == 1 { + return b[0], nil + } + if errors.Is(err, io.EOF) { + return 0, errInputEnded + } + if err != nil { + return 0, fmt.Errorf("read the terminal: %w", err) + } + + return 0, errInputEnded +} + +// Checklist is the live list of the steps of one job: a spinner on the step that runs, a mark on each one done. +type Checklist struct { + t *Terminal + mu sync.Mutex + steps []checkStep + drawn int + frame int + // spin stops the spinner; nil while no step runs. + spin chan struct{} + wg sync.WaitGroup + // err is the first failed write of the spinner, returned by the next call. + err error +} + +type checkState int + +const ( + pending checkState = iota + running + done + attention + failed +) + +type checkStep struct { + title string + state checkState + detail []string +} + +var spinner = []string{"⠋", "⠙", "⠹", "⠸", "⠼", "⠴", "⠦", "⠧", "⠇", "⠏"} + +const spinInterval = 80 * time.Millisecond + +func newTicker() *time.Ticker { return time.NewTicker(spinInterval) } + +// Checklist prints the title and every step as pending; off a terminal it prints each step once, as it ends. +func (t *Terminal) Checklist(title string, steps []string) (*Checklist, error) { + list := &Checklist{t: t} + for _, s := range steps { + list.steps = append(list.steps, checkStep{title: s}) + } + if err := t.Print(title, ""); err != nil { + return nil, err + } + if !t.interactive { + return list, nil + } + + return list, list.draw() +} + +// Start marks step i as the one that runs. +func (c *Checklist) Start(i int) error { + c.mu.Lock() + defer c.mu.Unlock() + c.steps[i].state = running + if !c.t.interactive { + return c.err + } + if c.spin == nil { + c.spin = make(chan struct{}) + c.wg.Add(1) + go c.turn(c.spin) + } + + return errors.Join(c.err, c.draw()) +} + +// Done marks step i complete. +func (c *Checklist) Done(i int) error { return c.end(i, done, nil) } + +// Attention marks step i complete with something the reader must see, in the lines under it. +func (c *Checklist) Attention(i int, detail ...string) error { return c.end(i, attention, detail) } + +// Fail marks step i failed, with the reason in the lines under it. +func (c *Checklist) Fail(i int, detail ...string) error { return c.end(i, failed, detail) } + +func (c *Checklist) end(i int, state checkState, detail []string) error { + c.stopSpinner() + + c.mu.Lock() + defer c.mu.Unlock() + c.steps[i].state, c.steps[i].detail = state, detail + if c.t.interactive { + return errors.Join(c.err, c.draw()) + } + + return errors.Join(c.err, c.t.Print(c.stepLines(c.steps[i])...)) +} + +func (c *Checklist) stopSpinner() { + c.mu.Lock() + spin := c.spin + c.spin = nil + c.mu.Unlock() + if spin != nil { + close(spin) + c.wg.Wait() + } +} + +func (c *Checklist) turn(stop chan struct{}) { + defer c.wg.Done() + tick := newTicker() + defer tick.Stop() + for { + select { + case <-stop: + return + case <-tick.C: + c.mu.Lock() + c.frame++ + if err := c.draw(); err != nil && c.err == nil { + c.err = err + } + c.mu.Unlock() + } + } +} + +// draw repaints the whole list in place; the caller holds the lock. +func (c *Checklist) draw() error { + var lines []string + for _, s := range c.steps { + lines = append(lines, c.stepLines(s)...) + } + drawn, err := c.t.redraw(c.drawn, lines, "\n") + c.drawn = drawn + + return err +} + +func (c *Checklist) stepLines(s checkStep) []string { + mark := map[checkState]string{ + pending: c.t.paint(gray, "○"), + running: c.t.paint(blue, spinner[c.frame%len(spinner)]), + done: c.t.paint(green, "✓"), + attention: c.t.paint(yellow, "!"), + failed: c.t.paint(red, "✗"), + }[s.state] + lines := []string{mark + " " + s.title} + for _, d := range s.detail { + lines = append(lines, " "+d) + } + + return lines +} diff --git a/pkg/term/term_test.go b/pkg/term/term_test.go new file mode 100644 index 000000000..16df73d6d --- /dev/null +++ b/pkg/term/term_test.go @@ -0,0 +1,149 @@ +package term + +import ( + "bytes" + "context" + "errors" + "io" + "os" + "strings" + "testing" + + "github.com/presmihaylov/shard/pkg/pty" +) + +// keyed is an interactive terminal without color whose keys are typed. +func keyed(out io.Writer, typed string) *Terminal { + keys := strings.NewReader(typed) + + return &Terminal{ + out: out, + interactive: true, + input: func(context.Context) io.Reader { return keys }, + raw: func() (pty.Restore, error) { return func() error { return nil }, nil }, + } +} + +func TestOffATerminalEveryQuestionNeedsOne(t *testing.T) { + in, w, err := os.Pipe() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + if err := errors.Join(in.Close(), w.Close()); err != nil { + t.Error(err) + } + }) + term := New(in, &bytes.Buffer{}, func(string) string { return "" }) + if term.Interactive() { + t.Fatal("a pipe is interactive") + } + + _, selectErr := term.Select(t.Context(), "pick", []Option{{Name: "a", Label: "A"}}) + _, confirmErr := term.Confirm(t.Context(), "sure?", true) + _, textErr := term.Text(t.Context(), "url") + _, secretErr := term.Secret(t.Context(), "key") + for name, err := range map[string]error{"select": selectErr, "confirm": confirmErr, "text": textErr, "secret": secretErr} { + if !errors.Is(err, ErrNotTerminal) { + t.Errorf("%s off a terminal: %v, want ErrNotTerminal", name, err) + } + } +} + +func TestOffATerminalTheChecklistPrintsEachStepOnceItEnds(t *testing.T) { + var out bytes.Buffer + list, err := (&Terminal{out: &out}).Checklist("Setting up", []string{"one", "two", "three"}) + if err != nil { + t.Fatal(err) + } + for _, mark := range []error{list.Start(0), list.Done(0), list.Start(1), list.Fail(1, "it broke")} { + if mark != nil { + t.Fatal(mark) + } + } + + if want := "Setting up\n\n✓ one\n✗ two\n it broke\n"; out.String() != want { + t.Errorf("printed %q, want %q", out.String(), want) + } +} + +func TestSelectMovesPastTheUnavailableRows(t *testing.T) { + options := []Option{ + {Name: "a", Label: "A", Default: true}, + {Name: "b", Label: "B", Unavailable: []string{"needs /dev/kvm"}}, + {Name: "c", Label: "C"}, + } + for _, tc := range []struct { + typed string + want int + }{ + {"\r", 0}, + {"\x1b[B\r", 2}, + {"jj\r", 2}, + {"j\x1b[Ak\r", 0}, + } { + var out bytes.Buffer + chosen, err := keyed(&out, tc.typed).Select(t.Context(), "pick", options) + if err != nil { + t.Fatalf("typed %q: %v", tc.typed, err) + } + if chosen != tc.want { + t.Errorf("typed %q chose %d, want %d", tc.typed, chosen, tc.want) + } + if !strings.Contains(out.String(), " B [Unavailable]") || !strings.Contains(out.String(), " needs /dev/kvm") { + t.Errorf("typed %q drew %q without the unavailable row and its reason", tc.typed, out.String()) + } + } +} + +func TestSelectRefusesWhenNoRowCanBeChosen(t *testing.T) { + _, err := keyed(&bytes.Buffer{}, "\r").Select(t.Context(), "pick", []Option{{Label: "A", Unavailable: []string{"no"}}}) + if err == nil { + t.Fatal("a select with every row unavailable chose one") + } +} + +func TestConfirmTakesTheDefaultOnAnEmptyReply(t *testing.T) { + for _, tc := range []struct { + typed string + yes bool + want bool + }{ + {"\n", true, true}, + {"\n", false, false}, + {"maybe\nn\n", true, false}, + {"YES\n", false, true}, + } { + got, err := keyed(&bytes.Buffer{}, tc.typed).Confirm(t.Context(), "sure?", tc.yes) + if err != nil { + t.Fatalf("typed %q: %v", tc.typed, err) + } + if got != tc.want { + t.Errorf("typed %q with default %v answered %v", tc.typed, tc.yes, got) + } + } +} + +func TestSecretEchoesOneDotPerCharacter(t *testing.T) { + var out bytes.Buffer + got, err := keyed(&out, "kéy\x7fy\x15ab\r").Secret(t.Context(), "API key") + if err != nil { + t.Fatal(err) + } + if got != "ab" { + t.Errorf("read %q, want ab", got) + } + echo, ok := strings.CutPrefix(out.String(), "API key\r\n> ") + if rest := strings.NewReplacer("•", "", "\b \b", "", "\r\n", "").Replace(echo); !ok || rest != "" { + t.Errorf("the secret echoed in %q", out.String()) + } + if dots := strings.Count(echo, "•"); dots != 6 { + t.Errorf("echoed %d dots, want 6", dots) + } +} + +func TestSecretEndsOnInterrupt(t *testing.T) { + if _, err := keyed(&bytes.Buffer{}, "ab\x03").Secret(t.Context(), "API key"); !errors.Is(err, ErrInterrupted) { + t.Errorf("Ctrl-C gave %v, want ErrInterrupted", err) + } +} diff --git a/services/client/client.go b/services/client/client.go index 107b80b33..988cd5817 100644 --- a/services/client/client.go +++ b/services/client/client.go @@ -234,6 +234,24 @@ func (c *Client) dial(ctx context.Context) (net.Conn, error) { return conn, nil } +// Reach opens one connection to the server and closes it, the tls handshake included, so a caller tells a server it cannot reach from one that refuses its token. +func (c *Client) Reach(ctx context.Context) error { + if c.Timeout != 0 { + var cancel context.CancelFunc + ctx, cancel = context.WithTimeout(ctx, c.Timeout) + defer cancel() + } + conn, err := c.dial(ctx) + if err != nil { + return err + } + if err := conn.Close(); err != nil { + return fmt.Errorf("close the connection to %s: %w", c.target, err) + } + + return nil +} + // authorize carries the bearer token of a front. The socket takes none: its mode is the check. func (c *Client) authorize(header http.Header) { if c.token == "" { diff --git a/services/client/config.go b/services/client/config.go new file mode 100644 index 000000000..759239c37 --- /dev/null +++ b/services/client/config.go @@ -0,0 +1,135 @@ +package client + +import ( + "bytes" + "encoding/json" + "errors" + "fmt" + "io/fs" + "os" + "os/user" + "path/filepath" + "strings" + + "github.com/presmihaylov/shard/pkg/store" +) + +// ConfigHomeEnv moves the user's configuration directory, as the XDG base directory spec defines it. +const ConfigHomeEnv = "XDG_CONFIG_HOME" + +// Config is the saved connection that shard setup writes; it holds the key as plain text, so only the user can read the file. +type Config struct { + Remote string `json:"remote"` + APIKey string `json:"api_key"` +} + +// Format prints the remote alone, its password hidden, so a config in an error or a log line never shows its key. +func (c Config) Format(f fmt.State, _ rune) { + fmt.Fprintf(f, "saved connection to %s", Redacted(c.Remote)) +} + +// ConfigPath is $XDG_CONFIG_HOME/shard/config.json, or ~/.config/shard/config.json; the XDG spec ignores a relative XDG_CONFIG_HOME. +func ConfigPath(getenv func(string) string) (string, error) { + if dir := getenv(ConfigHomeEnv); filepath.IsAbs(dir) { + return filepath.Join(dir, "shard", "config.json"), nil + } + + home := getenv("HOME") + if home == "" { + // A service manager may start a verb with no HOME, and the account database still names the home directory. + current, err := user.Current() + if err != nil { + return "", fmt.Errorf("find the shard configuration: HOME is not set and %w", err) + } + home = current.HomeDir + } + if home == "" { + return "", errors.New("find the shard configuration: HOME is not set and the user has no home directory") + } + + return filepath.Join(home, ".config", "shard", "config.json"), nil +} + +// LoadConfig reads the saved connection at path. No file is the zero Config, since a host with no saved connection is the common case. +func LoadConfig(path string) (Config, error) { + data, err := os.ReadFile(path) + if errors.Is(err, fs.ErrNotExist) { + return Config{}, nil + } + if err != nil { + return Config{}, fmt.Errorf("read the shard configuration %s: %w", path, err) + } + + var c Config + decoder := json.NewDecoder(bytes.NewReader(data)) + decoder.DisallowUnknownFields() + if err := decoder.Decode(&c); err != nil { + return Config{}, fmt.Errorf("parse the shard configuration %s: %w", path, err) + } + if err := c.check(); err != nil { + return Config{}, fmt.Errorf("the shard configuration %s %w", path, err) + } + + return c, nil +} + +// SaveConfig replaces the file at path with c whole, by a rename, in a directory only the user can open. +func SaveConfig(path string, c Config) error { + if c.Remote == "" || strings.TrimSpace(c.APIKey) == "" { + return errors.New("a saved connection needs a remote and an API key") + } + if err := c.check(); err != nil { + return fmt.Errorf("the connection %w", err) + } + + data, err := json.MarshalIndent(c, "", " ") //nolint:gosec // the file exists to hold the key, for its user alone + if err != nil { + return fmt.Errorf("encode the shard configuration: %w", err) + } + + dir := filepath.Dir(path) + if err := store.MkdirAllDurable(dir, 0o700); err != nil { + return fmt.Errorf("create the shard configuration directory: %w", err) + } + // MkdirAll leaves the mode of a directory that was already there. + if err := os.Chmod(dir, 0o700); err != nil { //nolint:gosec // a directory needs its search bit, and only its user has it + return fmt.Errorf("restrict %s to its user: %w", dir, err) + } + + if err := store.WriteFile(path, append(data, '\n'), 0o600); err != nil { + return fmt.Errorf("write the shard configuration: %w", err) + } + + return nil +} + +// RemoveConfig deletes the saved connection; a file that is already gone is not an error. +func RemoveConfig(path string) error { + err := os.Remove(path) + if errors.Is(err, fs.ErrNotExist) { + return nil + } + if err != nil { + return fmt.Errorf("remove the shard configuration %s: %w", path, err) + } + + if err := store.SyncDir(filepath.Dir(path)); err != nil { + return fmt.Errorf("remove the shard configuration %s: %w", path, err) + } + + return nil +} + +// check refuses a remote the client cannot dial and a key no header can carry; an error never quotes the key. +func (c Config) check() error { + if c.Remote != "" { + if _, err := parseRemote(c.Remote); err != nil { + return fmt.Errorf("holds a remote that %w", err) + } + } + if err := checkToken(c.APIKey); err != nil { + return fmt.Errorf("holds an API key that %w", err) + } + + return nil +} diff --git a/services/client/config_test.go b/services/client/config_test.go new file mode 100644 index 000000000..a3fa58dd4 --- /dev/null +++ b/services/client/config_test.go @@ -0,0 +1,214 @@ +package client_test + +import ( + "fmt" + "os" + "os/user" + "path/filepath" + "strings" + "testing" + + "github.com/presmihaylov/shard/services/client" +) + +// env answers the variables a test names and the empty string for every other, so the shell that runs the test decides nothing. +func env(vars map[string]string) func(string) string { + return func(name string) string { return vars[name] } +} + +// The file lives under XDG_CONFIG_HOME, or ~/.config, a relative XDG_CONFIG_HOME is ignored as the XDG spec says, and no HOME asks the account database. (SHARD-657) +func TestConfigPathFollowsTheXDGSpec(t *testing.T) { + for _, tc := range []struct { + name string + vars map[string]string + want string + }{ + {name: "xdg", vars: map[string]string{client.ConfigHomeEnv: "/xdg", "HOME": "/home/u"}, want: "/xdg/shard/config.json"}, + {name: "no xdg", vars: map[string]string{"HOME": "/home/u"}, want: "/home/u/.config/shard/config.json"}, + {name: "relative xdg", vars: map[string]string{client.ConfigHomeEnv: "xdg", "HOME": "/home/u"}, want: "/home/u/.config/shard/config.json"}, + } { + t.Run(tc.name, func(t *testing.T) { + got, err := client.ConfigPath(env(tc.vars)) + if err != nil { + t.Fatalf("ConfigPath: %v", err) + } + if got != tc.want { + t.Errorf("ConfigPath answered %s, want %s", got, tc.want) + } + }) + } + + current, err := user.Current() + if err != nil { + t.Fatalf("user.Current: %v", err) + } + got, err := client.ConfigPath(env(nil)) + if err != nil { + t.Fatalf("ConfigPath with no HOME: %v", err) + } + if want := filepath.Join(current.HomeDir, ".config", "shard", "config.json"); got != want { + t.Errorf("ConfigPath with no HOME answered %s, want %s from the account database", got, want) + } +} + +// A saved connection is the user's alone, the directory 0700 and the file 0600, even where the directory was wider. (SHARD-657) +func TestSaveConfigRestrictsTheFileAndItsDirectory(t *testing.T) { + dir := filepath.Join(t.TempDir(), "shard") + if err := os.Mkdir(dir, 0o755); err != nil { + t.Fatalf("make the directory: %v", err) + } + path := filepath.Join(dir, "config.json") + + saved := client.Config{Remote: "https://shard.example.com", APIKey: leakKey} + if err := client.SaveConfig(path, saved); err != nil { + t.Fatalf("SaveConfig: %v", err) + } + + for p, want := range map[string]os.FileMode{dir: 0o700, path: 0o600} { + info, err := os.Stat(p) + if err != nil { + t.Fatalf("stat %s: %v", p, err) + } + if got := info.Mode().Perm(); got != want { + t.Errorf("%s has mode %o, want %o", p, got, want) + } + } + + got, err := client.LoadConfig(path) + if err != nil { + t.Fatalf("LoadConfig: %v", err) + } + if got != saved { + t.Errorf("LoadConfig answered a different connection than SaveConfig wrote") + } +} + +// A save replaces the whole file through a rename, and leaves no temp file beside it. (SHARD-657) +func TestSaveConfigReplacesTheWholeFile(t *testing.T) { + path := filepath.Join(t.TempDir(), "shard", "config.json") + + if err := client.SaveConfig(path, client.Config{Remote: "https://old.example.com", APIKey: "old-key"}); err != nil { + t.Fatalf("SaveConfig the old connection: %v", err) + } + replacement := client.Config{Remote: "http://new.example.com:2376", APIKey: "new-key"} + if err := client.SaveConfig(path, replacement); err != nil { + t.Fatalf("SaveConfig the new connection: %v", err) + } + + got, err := client.LoadConfig(path) + if err != nil { + t.Fatalf("LoadConfig: %v", err) + } + if got != replacement { + t.Errorf("LoadConfig answered %v, want the replacement", got) + } + + entries, err := os.ReadDir(filepath.Dir(path)) + if err != nil { + t.Fatalf("read the directory: %v", err) + } + if len(entries) != 1 { + t.Errorf("the directory holds %d entries after two saves, want config.json alone", len(entries)) + } +} + +// A connection that is incomplete or that no client can use is never written, and the refusal never quotes the key. (SHARD-657) +func TestSaveConfigRefusesAConnectionNoClientCanUse(t *testing.T) { + path := filepath.Join(t.TempDir(), "shard", "config.json") + + for name, c := range map[string]client.Config{ + "no remote": {APIKey: leakKey}, + "no key": {Remote: "https://shard.example.com"}, + "blank key": {Remote: "https://shard.example.com", APIKey: " \t"}, + "ftp remote": {Remote: "ftp://shard.example.com", APIKey: leakKey}, + "split key": {Remote: "https://shard.example.com", APIKey: leakKey + "\nX-Other: 1"}, + } { + t.Run(name, func(t *testing.T) { + err := client.SaveConfig(path, c) + if err == nil { + t.Fatal("SaveConfig wrote the connection") + } + if strings.Contains(err.Error(), leakKey) { + t.Errorf("the refusal %q quotes the key", err) + } + if _, err := os.Stat(path); !os.IsNotExist(err) { + t.Errorf("stat %s returned %v after a refused save, want no file", path, err) + } + }) + } +} + +// No file is no saved connection, and no error: most hosts save none. (SHARD-657) +func TestLoadConfigWithNoFileIsNoConnection(t *testing.T) { + got, err := client.LoadConfig(filepath.Join(t.TempDir(), "shard", "config.json")) + if err != nil { + t.Fatalf("LoadConfig: %v", err) + } + if got != (client.Config{}) { + t.Errorf("LoadConfig answered %v, want the zero connection", got) + } +} + +// A file the client cannot use is refused with its path, never read in part, and the refusal never quotes the key. (SHARD-657) +func TestLoadConfigRefusesAFileTheClientCannotUse(t *testing.T) { + for name, body := range map[string]string{ + "not json": "remote = https://shard.example.com", + "unknown field": `{"remote":"https://shard.example.com","api_key":"` + leakKey + `","token":"x"}`, + "ftp remote": `{"remote":"ftp://shard.example.com","api_key":"` + leakKey + `"}`, + "split key": `{"remote":"https://shard.example.com","api_key":"` + leakKey + `\r\nX-Other: 1"}`, + } { + t.Run(name, func(t *testing.T) { + path := filepath.Join(t.TempDir(), "config.json") + if err := os.WriteFile(path, []byte(body), 0o600); err != nil { + t.Fatalf("write the file: %v", err) + } + + _, err := client.LoadConfig(path) + if err == nil { + t.Fatal("LoadConfig read the file") + } + if !strings.Contains(err.Error(), path) { + t.Errorf("the refusal %q does not name the file", err) + } + if strings.Contains(err.Error(), leakKey) { + t.Errorf("the refusal %q quotes the key", err) + } + }) + } +} + +// Remove deletes the saved connection, and one already gone is not an error. (SHARD-657) +func TestRemoveConfig(t *testing.T) { + path := filepath.Join(t.TempDir(), "shard", "config.json") + if err := client.SaveConfig(path, client.Config{Remote: "https://shard.example.com", APIKey: leakKey}); err != nil { + t.Fatalf("SaveConfig: %v", err) + } + + for range 2 { + if err := client.RemoveConfig(path); err != nil { + t.Fatalf("RemoveConfig: %v", err) + } + } + if _, err := os.Stat(path); !os.IsNotExist(err) { + t.Errorf("stat %s returned %v after RemoveConfig, want no file", path, err) + } +} + +// A saved connection prints its remote alone, whatever the verb and with its password hidden, so no error or log line shows a secret. (SHARD-657) +func TestAConfigNeverPrintsItsKey(t *testing.T) { + c := client.Config{Remote: "https://shard.example.com", APIKey: leakKey} + + for _, verb := range []string{"%v", "%+v", "%#v", "%s", "%q"} { + got := fmt.Sprintf(verb, c) + if strings.Contains(got, leakKey) { + t.Errorf("%s printed the key: %s", verb, got) + } + if !strings.Contains(got, c.Remote) { + t.Errorf("%s printed %q, want the remote", verb, got) + } + } + + if got := fmt.Sprint(client.Config{Remote: "https://user:" + leakKey + "@shard.example.com", APIKey: "saved-key"}); strings.Contains(got, leakKey) { + t.Errorf("a remote with a password printed %q", got) + } +} diff --git a/services/client/credential.go b/services/client/credential.go index 9ecfb5718..163195a88 100644 --- a/services/client/credential.go +++ b/services/client/credential.go @@ -4,6 +4,7 @@ import ( "cmp" "errors" "fmt" + "net/url" "os" "strings" ) @@ -15,25 +16,38 @@ const ( CAFileEnv = "SHARD_CA_FILE" ) -// NewRemoteFromEnv builds the client of host, or of SHARD_REMOTE when host is empty, with SHARD_API_KEY and SHARD_CA_FILE. -func NewRemoteFromEnv(host string) (*Client, error) { - host = cmp.Or(host, os.Getenv(RemoteEnv)) +// NewRemoteFromEnv builds the client of host, or of SHARD_REMOTE, or of the saved remote, with SHARD_API_KEY or the saved key, and SHARD_CA_FILE. +func NewRemoteFromEnv(host string, saved Config) (*Client, error) { + host = cmp.Or(host, os.Getenv(RemoteEnv), saved.Remote) if host == "" { - return nil, fmt.Errorf("a remote client needs a host: --remote or %s, as https://shard.example.com", RemoteEnv) + return nil, fmt.Errorf("a remote client needs a host: --remote, %s or shard setup, as https://shard.example.com", RemoteEnv) } parsed, err := parseRemote(host) if err != nil { return nil, err } - token, err := apiKey() + token, err := apiKey(parsed, saved) if err != nil { return nil, err } - caFile := os.Getenv(CAFileEnv) + ca, err := ReadCA(host, os.Getenv(CAFileEnv)) + if err != nil { + return nil, err + } + + return NewRemote(host, token, ca) +} + +// ReadCA reads the certificate caFile, SHARD_CA_FILE, holds for host, or none when it is empty. +func ReadCA(host, caFile string) ([]byte, error) { if caFile == "" { - return NewRemote(host, token, nil) + return nil, nil + } + parsed, err := parseRemote(host) + if err != nil { + return nil, err } // Refused before the file is read, so the mismatch is the one error, whatever the file holds. if parsed.Scheme == "http" { @@ -44,20 +58,49 @@ func NewRemoteFromEnv(host string) (*Client, error) { return nil, fmt.Errorf("read the ca file %s from %s: %w", caFile, CAFileEnv, err) } - return NewRemote(host, token, ca) + return ca, nil } -// apiKey is SHARD_API_KEY, the one credential of a remote client; an error names the variable, never the value. -func apiKey() (string, error) { +// apiKey is SHARD_API_KEY, else the saved key when remote is the saved remote, so a key reaches no server it was not saved for; an error never quotes it. +func apiKey(remote *url.URL, saved Config) (string, error) { key := strings.TrimSpace(os.Getenv(APIKeyEnv)) - if key == "" { - return "", fmt.Errorf("a remote client needs %s; shard serve answers 401 without one", APIKeyEnv) + if key != "" { + if err := checkToken(key); err != nil { + return "", fmt.Errorf("%s %w", APIKeyEnv, err) + } + + return key, nil } - if err := checkToken(key); err != nil { - return "", fmt.Errorf("%s %w", APIKeyEnv, err) + + key = strings.TrimSpace(saved.APIKey) + if key != "" && sameRemote(remote, saved.Remote) { + return key, nil + } + if key != "" { + return "", fmt.Errorf("a remote client needs %s for %s; the saved API key is for %s and is sent to no other server", APIKeyEnv, remote.Redacted(), Redacted(saved.Remote)) + } + + return "", fmt.Errorf("a remote client needs %s; shard serve answers 401 without one", APIKeyEnv) +} + +// sameRemote compares what the client dials, the scheme and the host and port, so a path or the case of a host name makes no other server. +func sameRemote(remote *url.URL, saved string) bool { + other, err := url.Parse(saved) + if err != nil { + return false + } + + return remote.Scheme == other.Scheme && strings.EqualFold(remoteAddress(remote), remoteAddress(other)) +} + +// Redacted is remote with the password it may carry hidden, as url.URL.Redacted prints it. +func Redacted(remote string) string { + parsed, err := url.Parse(remote) + if err != nil { + return remote } - return key, nil + return parsed.Redacted() } // checkToken refuses the bytes net/http refuses in a header, so a newline never splits a request; any other wrong token is the front's to refuse. diff --git a/services/client/credential_test.go b/services/client/credential_test.go index 5e04f99f0..409a044dd 100644 --- a/services/client/credential_test.go +++ b/services/client/credential_test.go @@ -1,7 +1,9 @@ package client_test import ( + "crypto/tls" "encoding/pem" + "errors" "fmt" "net/http" "net/http/httptest" @@ -88,7 +90,7 @@ func TestNewRemoteFromEnvSendsTheKeyOverEitherScheme(t *testing.T) { t.Setenv(client.APIKeyEnv, "key-token") t.Setenv(client.CAFileEnv, tc.ca) - c, err := client.NewRemoteFromEnv("") + c, err := client.NewRemoteFromEnv("", client.Config{}) if err != nil { t.Fatalf("NewRemoteFromEnv: %v", err) } @@ -112,7 +114,7 @@ func TestAMissingKeyNamesSHARDAPIKEYAlone(t *testing.T) { noRemoteEnv(t) t.Setenv(client.APIKeyEnv, key) - _, err := client.NewRemoteFromEnv("https://shard.example.com") + _, err := client.NewRemoteFromEnv("https://shard.example.com", client.Config{}) if err == nil { t.Fatal("NewRemoteFromEnv answered with no key") } @@ -129,7 +131,7 @@ func TestHTTPSTrustsTheStoreOrSHARDCAFILE(t *testing.T) { noRemoteEnv(t) t.Setenv(client.APIKeyEnv, "key-token") - c, err := client.NewRemoteFromEnv(host) + c, err := client.NewRemoteFromEnv(host, client.Config{}) if err != nil { t.Fatalf("NewRemoteFromEnv: %v", err) } @@ -138,7 +140,7 @@ func TestHTTPSTrustsTheStoreOrSHARDCAFILE(t *testing.T) { } t.Setenv(client.CAFileEnv, ca) - c, err = client.NewRemoteFromEnv(host) + c, err = client.NewRemoteFromEnv(host, client.Config{}) if err != nil { t.Fatalf("NewRemoteFromEnv: %v", err) } @@ -158,7 +160,7 @@ func TestSHARDCAFILEWithAnHTTPRemoteIsRefused(t *testing.T) { t.Setenv(client.APIKeyEnv, "key-token") t.Setenv(client.CAFileEnv, caFile) - _, err := client.NewRemoteFromEnv(host) + _, err := client.NewRemoteFromEnv(host, client.Config{}) if err == nil || !strings.Contains(err.Error(), client.CAFileEnv) || !strings.Contains(err.Error(), host) { t.Errorf("%s with %s returned %v, want a refusal that names both", client.CAFileEnv, host, err) } @@ -179,7 +181,7 @@ func TestNewRemoteFromEnvNeedsAHost(t *testing.T) { noRemoteEnv(t) t.Setenv(client.APIKeyEnv, "key-token") - _, err := client.NewRemoteFromEnv("") + _, err := client.NewRemoteFromEnv("", client.Config{}) if err == nil || !strings.Contains(err.Error(), client.RemoteEnv) { t.Errorf("NewRemoteFromEnv with no host returned %v, want a refusal that names %s", err, client.RemoteEnv) } @@ -215,7 +217,7 @@ func TestNoErrorHoldsTheKey(t *testing.T) { t.Setenv(client.APIKeyEnv, tc.key) t.Setenv(client.CAFileEnv, tc.caFile) - _, err := client.NewRemoteFromEnv("") + _, err := client.NewRemoteFromEnv("", client.Config{}) if err == nil { t.Fatal("NewRemoteFromEnv accepted it") } @@ -246,7 +248,7 @@ func TestNoErrorHoldsTheKey(t *testing.T) { t.Setenv(client.APIKeyEnv, key) t.Setenv(client.CAFileEnv, ca) - c, err := client.NewRemoteFromEnv("") + c, err := client.NewRemoteFromEnv("", client.Config{}) if err != nil { t.Fatalf("NewRemoteFromEnv: %v", err) } @@ -263,3 +265,186 @@ func TestNoErrorHoldsTheKey(t *testing.T) { }) } } + +// sawRequest says whether the front saw a request, without waiting for one. +func sawRequest(seen chan string) bool { + select { + case <-seen: + return true + default: + return false + } +} + +// The remote is --remote, else SHARD_REMOTE, else the saved one, and only the chosen server sees a request. (SHARD-657) +func TestTheRemoteIsTheFlagThenSHARDREMOTEThenTheSavedOne(t *testing.T) { + flag, flagSeen := plainFront(t, "key-token") + envRemote, envSeen := plainFront(t, "key-token") + saved, savedSeen := plainFront(t, "key-token") + + for _, tc := range []struct { + name, host, env string + want chan string + }{ + {name: "the flag", host: flag, env: envRemote, want: flagSeen}, + {name: "SHARD_REMOTE", env: envRemote, want: envSeen}, + {name: "the saved remote", want: savedSeen}, + } { + t.Run(tc.name, func(t *testing.T) { + noRemoteEnv(t) + t.Setenv(client.RemoteEnv, tc.env) + t.Setenv(client.APIKeyEnv, "key-token") + + c, err := client.NewRemoteFromEnv(tc.host, client.Config{Remote: saved, APIKey: "key-token"}) + if err != nil { + t.Fatalf("NewRemoteFromEnv: %v", err) + } + if _, err := c.Version(t.Context()); err != nil { + t.Fatalf("Version: %v", err) + } + for name, seen := range map[string]chan string{"the flag": flagSeen, "SHARD_REMOTE": envSeen, "the saved remote": savedSeen} { + if got := sawRequest(seen); got != (seen == tc.want) { + t.Errorf("the front of %s saw a request: %v", name, got) + } + } + }) + } +} + +// SHARD_API_KEY beats the saved key, and an empty or blank one counts as unset, so the saved key rides. (SHARD-657) +func TestSHARDAPIKEYBeatsTheSavedKey(t *testing.T) { + for _, tc := range []struct { + name, env, want string + }{ + {name: "set", env: "env-key", want: "env-key"}, + {name: "unset", env: "", want: "saved-key"}, + {name: "blank", env: " \t\n", want: "saved-key"}, + } { + t.Run(tc.name, func(t *testing.T) { + host, seen := plainFront(t, tc.want) + noRemoteEnv(t) + t.Setenv(client.APIKeyEnv, tc.env) + + c, err := client.NewRemoteFromEnv("", client.Config{Remote: host, APIKey: "saved-key"}) + if err != nil { + t.Fatalf("NewRemoteFromEnv: %v", err) + } + if _, err := c.Version(t.Context()); err != nil { + t.Fatalf("Version: %v", err) + } + if got := <-seen; got != "Bearer "+tc.want { + t.Errorf("the front saw %q, want %q as the bearer", got, tc.want) + } + }) + } +} + +// The saved key goes to the server it was saved for alone; another remote needs SHARD_API_KEY, and the refusal never quotes the key. (SHARD-657) +func TestTheSavedKeyGoesToTheSavedRemoteAlone(t *testing.T) { + other, otherSeen := plainFront(t, leakKey) + saved, _ := plainFront(t, leakKey) + + for name, set := range map[string]func(*testing.T) string{ + "the flag": func(*testing.T) string { return other }, + "SHARD_REMOTE": func(t *testing.T) string { t.Setenv(client.RemoteEnv, other); return "" }, + } { + t.Run(name, func(t *testing.T) { + noRemoteEnv(t) + host := set(t) + + _, err := client.NewRemoteFromEnv(host, client.Config{Remote: saved, APIKey: leakKey}) + if err == nil { + t.Fatal("NewRemoteFromEnv sent the saved key to another server") + } + if msg := err.Error(); !strings.Contains(msg, client.APIKeyEnv) || !strings.Contains(msg, saved) || strings.Contains(msg, leakKey) { + t.Errorf("NewRemoteFromEnv returned %q, want %s and the saved remote in it and never the key", msg, client.APIKeyEnv) + } + if sawRequest(otherSeen) { + t.Error("the other server saw a request") + } + }) + } +} + +// One server is one scheme and one host and port: the case of the name, a path and the default port make no other, a scheme or a port does. (SHARD-657) +func TestTheSavedKeyMatchesItsServerWhateverTheSpelling(t *testing.T) { + saved := client.Config{Remote: "https://Shard.Example.com/", APIKey: leakKey} + + for host, same := range map[string]bool{ + "https://shard.example.com": true, + "https://shard.example.com:443/v0": true, + "http://shard.example.com": false, + "https://shard.example.com:8443": false, + "https://shard.example.com.evil.io": false, + } { + t.Run(host, func(t *testing.T) { + noRemoteEnv(t) + + _, err := client.NewRemoteFromEnv(host, saved) + if same && err != nil { + t.Errorf("NewRemoteFromEnv refused the saved server: %v", err) + } + if !same && err == nil { + t.Error("NewRemoteFromEnv sent the saved key to another server") + } + }) + } +} + +// A refusal names the saved remote with its password hidden, as url.URL.Redacted prints it. (SHARD-657) +func TestARefusalHidesThePasswordOfTheSavedRemote(t *testing.T) { + noRemoteEnv(t) + + _, err := client.NewRemoteFromEnv("https://other.example.com", client.Config{Remote: "https://user:" + leakKey + "@shard.example.com", APIKey: "saved-key"}) + if err == nil || strings.Contains(err.Error(), leakKey) || !strings.Contains(err.Error(), "shard.example.com") { + t.Errorf("NewRemoteFromEnv returned %v, want the saved remote in it and never its password", err) + } +} + +// Reach dials and shakes hands and sends no request, so a server that refuses the key still answers it, and an untrusted certificate fails it. (SHARD-657) +func TestReachDialsWithoutARequest(t *testing.T) { + host, ca, seen := tlsFront(t, "key-token") + caBytes, err := client.ReadCA(host, ca) + if err != nil { + t.Fatalf("ReadCA: %v", err) + } + + c, err := client.NewRemote(host, "not-the-key", caBytes) + if err != nil { + t.Fatalf("NewRemote: %v", err) + } + if err := c.Reach(t.Context()); err != nil { + t.Errorf("Reach with a key the front refuses: %v", err) + } + if sawRequest(seen) { + t.Error("Reach sent a request") + } + + c, err = client.NewRemote(host, "key-token", nil) + if err != nil { + t.Fatalf("NewRemote: %v", err) + } + var untrusted *tls.CertificateVerificationError + if err := c.Reach(t.Context()); !errors.As(err, &untrusted) { + t.Errorf("Reach against a private CA with none returned %v, want an untrusted certificate", err) + } + + listener := httptest.NewServer(http.NotFoundHandler()) + listener.Close() + c, err = client.NewRemote(listener.URL, "key-token", nil) + if err != nil { + t.Fatalf("NewRemote: %v", err) + } + var connect *client.ConnectError + if err := c.Reach(t.Context()); !errors.As(err, &connect) { + t.Errorf("Reach of a closed port returned %v, want a ConnectError", err) + } +} + +// No SHARD_CA_FILE is no certificate and no error. (SHARD-657) +func TestReadCAWithNoFileIsNone(t *testing.T) { + ca, err := client.ReadCA("https://shard.example.com", "") + if err != nil || ca != nil { + t.Errorf("ReadCA with no file answered %d bytes and %v, want none", len(ca), err) + } +} diff --git a/services/setup/apply_local.go b/services/setup/apply_local.go new file mode 100644 index 000000000..6e4b034e3 --- /dev/null +++ b/services/setup/apply_local.go @@ -0,0 +1,636 @@ +package setup + +import ( + "compress/gzip" + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "fmt" + "io" + "io/fs" + "maps" + "net/http" + "os" + "path" + "path/filepath" + "slices" + "strings" + "time" + + "github.com/klauspost/compress/zstd" + + "github.com/presmihaylov/shard/pkg/tarball" +) + +// The paths a local setup installs; a root service only ever runs root-owned files from them. +const ( + binDir = "/usr/local/bin" + shardBinary = binDir + "/shard" + initBinary = binDir + "/shard-init" + systemdUnit = "/etc/systemd/system/shard.service" + launchdPlist = "/Library/LaunchDaemons/shard.daemon.plist" + newsyslog = "/etc/newsyslog.d/shard.conf" + macLogDir = "/var/log/shard" + launchdLabel = "system/shard.daemon" +) + +// The bounds on one pinned archive, far above what any of them holds. +const ( + maxArchiveBytes = 1 << 30 + maxArchiveEntries = 256 +) + +// verifyWait covers a first Firecracker start, which fetches the guest kernel and formats the data image. +var ( + verifyWait = 2 * time.Minute + verifyPoll = time.Second +) + +// rooted is p on the host setup works on, which a test moves under a temp dir. +func rooted(h Host, p string) string { return filepath.Join(h.Root, p) } + +// privileged runs one command as root: as it is when setup runs as root, else under sudo, which admin has already let in. +func privileged(ctx context.Context, h Host, name string, args ...string) ([]byte, error) { + if h.Euid == 0 { + return run(ctx, h, name, args...) + } + + return run(ctx, h, "sudo", append([]string{"-n", "--", name}, args...)...) +} + +// run is one command whose failure carries what it printed. +func run(ctx context.Context, h Host, name string, args ...string) ([]byte, error) { + out, err := h.Run(ctx, name, args...) + if err != nil { + return out, fmt.Errorf("%s: %w%s", strings.Join(append([]string{name}, args...), " "), err, outputTail(out)) + } + + return out, nil +} + +// outputTail is the last line a command printed, which is where a tool says why it failed. +func outputTail(out []byte) string { + lines := strings.Split(strings.TrimSpace(string(out)), "\n") + last := strings.TrimSpace(lines[len(lines)-1]) + if last == "" { + return "" + } + + return ": " + last +} + +// admin makes sure the privileged steps that follow run: sudo asks for the password once, on the terminal, before any checklist draws. +func (s *Setup) admin(ctx context.Context) error { + if s.Host.Euid == 0 { + return nil + } + if _, err := s.Host.Run(ctx, "sudo", "-n", "true"); err == nil { + return nil + } + if err := s.UI.Print("", "Setup needs administrator access. sudo may ask for your password."); err != nil { + return err + } + if _, err := run(ctx, s.Host, "sudo", "-v"); err != nil { + return &Problem{Lines: []string{ + fmt.Sprintf("Administrator access failed: %v.", err), + "Run shard setup in a terminal where sudo can ask for your password, or as root.", + }} + } + + return nil +} + +// localPlan is one local setup worked out against the host before any change, which its steps then carry out. +type localPlan struct { + h Host + local Local + missing missing + // initAsset is the release file shard-init comes from, and "" on a Mac, where the daemon embeds it. + initAsset string + // user runs the daemon on a Mac and owns its root there. + user string + // stage holds the verified downloads between the download step and the install step. + stage string +} + +func newLocalPlan(h Host, l Local) *localPlan { + p := &localPlan{h: h, local: l, missing: missingTools(h, l.Provider)} + if h.OS == "linux" { + p.initAsset = "shard-init-linux-" + h.Arch + } + if h.OS == "darwin" { + p.user = daemonUser(h) + } + + return p +} + +// daemonUser is the person setup runs for, also when they started it with sudo. +func daemonUser(h Host) string { + if h.Euid == 0 { + return h.Env("SUDO_USER") + } + + return h.Env("USER") +} + +// localSteps are the steps that set up l on this host, in checklist order; each is safe to run again. +func (s *Setup) localSteps(_ context.Context, l Local) ([]Step, error) { + p := newLocalPlan(s.Host, l) + steps := []Step{{Title: "Check host compatibility", Do: p.compatible}} + if len(p.missing.downloads) > 0 || p.initAsset != "" { + steps = append(steps, Step{Title: "Download and verify required tools", Do: p.download}) + } + steps = append(steps, + Step{Title: "Install " + providerTitle(l.Provider), Do: p.install}, + Step{Title: "Create Shard's data directory", Do: p.dataDir}, + ) + if !l.StartAtBoot { + return steps, nil + } + + return append(steps, + Step{Title: "Configure the background service", Do: p.service}, + Step{Title: "Start the daemon", Do: p.start}, + Step{Title: "Verify the daemon connection", Do: p.verify}, + ), nil +} + +// compatible repeats the checks a change could have undone since preflight: the provider still runs here, and root is still in reach. +func (p *localPlan) compatible(ctx context.Context) error { + if l, ok := lacks(ctx, p.h, p.local.Provider); ok { + return &Problem{Lines: []string{providerTitle(p.local.Provider) + " requires " + l.need + ".", l.fact}} + } + if _, err := privileged(ctx, p.h, "true"); err != nil { + return &Problem{Lines: []string{fmt.Sprintf("Administrator access failed: %v.", err)}} + } + + return nil +} + +func (p *localPlan) download(ctx context.Context) (err error) { + stage, err := os.MkdirTemp("", "shard-setup-") + if err != nil { + return fmt.Errorf("create a download directory: %w", err) + } + p.stage = stage + defer func() { + if err != nil { + err = errors.Join(err, p.clean()) + } + }() + + for _, d := range p.missing.downloads { + if err := fetchPinned(ctx, p.h, d, stage); err != nil { + return &Problem{Lines: []string{fmt.Sprintf("Could not download %s: %v.", d.Title, err)}} + } + } + if p.initAsset == "" { + return nil + } + if err := FetchAsset(ctx, p.h, p.h.Version, p.initAsset, filepath.Join(stage, "shard-init"), 0o755); err != nil { + return &Problem{Lines: []string{fmt.Sprintf("Could not download shard-init: %v.", err)}} + } + + return nil +} + +func (p *localPlan) clean() error { + if p.stage == "" { + return nil + } + stage := p.stage + p.stage = "" + if err := os.RemoveAll(stage); err != nil { + return fmt.Errorf("remove the download directory: %w", err) + } + + return nil +} + +// fetchPinned downloads d into stage and refuses it unless it hashes to its pin; an archive is unpacked under stage/. +func fetchPinned(ctx context.Context, h Host, d *download, stage string) (err error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, d.URL, nil) + if err != nil { + return err + } + resp, err := h.HTTP.Do(req) + if err != nil { + return err + } + defer func() { err = errors.Join(err, resp.Body.Close()) }() + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("%s answered %s", d.URL, resp.Status) + } + + file := filepath.Join(stage, path.Base(d.URL)) + if err := saveVerified(resp.Body, file, d.SHA256); err != nil { + return err + } + if d.deb() { + return nil + } + + return unpack(file, filepath.Join(stage, d.Title)) +} + +func saveVerified(r io.Reader, file, want string) (err error) { + f, err := os.OpenFile(file, os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o600) + if err != nil { + return err + } + defer func() { err = errors.Join(err, f.Close()) }() + + sum := sha256.New() + if _, err := io.Copy(io.MultiWriter(f, sum), io.LimitReader(r, maxArchiveBytes)); err != nil { + return fmt.Errorf("save %s: %w", filepath.Base(file), err) + } + if got := hex.EncodeToString(sum.Sum(nil)); got != want { + return fmt.Errorf("its sha256 is %s, and setup expects %s", got, want) + } + + return nil +} + +// unpack extracts a .tgz or a .tar.zstd, which is how the pinned releases ship. +func unpack(file, dst string) (err error) { + f, err := os.Open(file) + if err != nil { + return err + } + defer func() { err = errors.Join(err, f.Close()) }() + + var tar io.Reader + switch { + case strings.HasSuffix(file, ".tgz"): + gz, gzErr := gzip.NewReader(f) + if gzErr != nil { + return fmt.Errorf("read %s: %w", filepath.Base(file), gzErr) + } + defer func() { err = errors.Join(err, gz.Close()) }() + tar = gz + case strings.HasSuffix(file, ".tar.zstd"): + zr, zErr := zstd.NewReader(f) + if zErr != nil { + return fmt.Errorf("read %s: %w", filepath.Base(file), zErr) + } + defer zr.Close() + tar = zr + default: + return fmt.Errorf("%s is no archive setup can unpack", filepath.Base(file)) + } + + if err := os.Mkdir(dst, 0o700); err != nil { + return err + } + if err := tarball.Unpack(tar, dst, tarball.Options{ConfineLinks: true, MaxBytes: maxArchiveBytes, MaxEntries: maxArchiveEntries}); err != nil { + return fmt.Errorf("unpack %s: %w", filepath.Base(file), err) + } + + return nil +} + +func (p *localPlan) install(ctx context.Context) (err error) { + defer func() { err = errors.Join(err, p.clean()) }() + + var owned []Owned + if p.missing.apt() { + if err := p.aptInstall(ctx); err != nil { + return err + } + } + for _, d := range p.missing.downloads { + for _, entry := range slices.Sorted(maps.Keys(d.Files)) { + dst := d.Files[entry] + if err := installFile(ctx, p.h, filepath.Join(p.stage, d.Title, entry), dst, "0755"); err != nil { + return &Problem{Lines: []string{fmt.Sprintf("Could not install %s: %v.", path.Base(dst), err)}} + } + owned = append(owned, Owned{Path: dst, Kind: KindTool}) + } + } + for _, t := range p.missing.tools { + if t.Download != nil && !t.Download.deb() { + continue + } + found, ok := lookPath(p.h, t.Name) + if !ok { + return &Problem{Lines: []string{fmt.Sprintf("%s is still missing after its package installed.", t.Name)}} + } + owned = append(owned, Owned{Path: found, Kind: KindTool, Package: toolPackage(t)}) + } + + binaries, err := p.installShard(ctx) + if err != nil { + return err + } + + return RecordOwned(ctx, p.h, p.local, append(owned, binaries...)...) +} + +// toolPackage is the apt package uninstall names for a tool setup installed. +func toolPackage(t tool) string { + if t.Download != nil && t.Download.deb() { + return strings.TrimSuffix(strings.Split(path.Base(t.Download.URL), "_")[0], ".deb") + } + + return t.Package +} + +func (p *localPlan) aptInstall(ctx context.Context) error { + if _, err := privileged(ctx, p.h, "apt-get", "update"); err != nil { + return &Problem{Lines: []string{fmt.Sprintf("Could not update the package lists: %v.", err)}} + } + args := []string{"DEBIAN_FRONTEND=noninteractive", "apt-get", "install", "-y", "--no-install-recommends"} + args = append(args, p.missing.packages...) + for _, d := range p.missing.downloads { + if d.deb() { + args = append(args, filepath.Join(p.stage, path.Base(d.URL))) + } + } + if _, err := privileged(ctx, p.h, "env", args...); err != nil { + return &Problem{Lines: []string{fmt.Sprintf("Could not install %s: %v.", strings.Join(p.missing.names(), ", "), err)}} + } + + return nil +} + +// installShard puts shard, and shard-init on Linux, where the service runs them: root-owned, since a root service never runs a file its user can change. +func (p *localPlan) installShard(ctx context.Context) ([]Owned, error) { + var owned []Owned + if p.h.Executable != shardBinary { + if err := installFile(ctx, p.h, p.h.Executable, shardBinary, "0755"); err != nil { + return nil, &Problem{Lines: []string{fmt.Sprintf("Could not install shard in %s: %v.", binDir, err)}} + } + owned = append(owned, Owned{Path: shardBinary, Kind: KindBinary}) + } + if p.initAsset == "" { + return owned, nil + } + if err := installFile(ctx, p.h, filepath.Join(p.stage, "shard-init"), initBinary, "0755"); err != nil { + return nil, &Problem{Lines: []string{fmt.Sprintf("Could not install shard-init in %s: %v.", binDir, err)}} + } + + return append(owned, Owned{Path: initBinary, Kind: KindBinary}), nil +} + +// installFile copies src to dst as root, with the directory above it; install replaces a file in use without writing into it. +func installFile(ctx context.Context, h Host, src, dst, mode string) error { + if _, err := privileged(ctx, h, "install", "-d", "-o", "0", "-g", "0", "-m", "0755", rooted(h, path.Dir(dst))); err != nil { + return err + } + _, err := privileged(ctx, h, "install", "-o", "0", "-g", "0", "-m", mode, src, rooted(h, dst)) + + return err +} + +// dataDir makes the root the daemon serves, and leaves one that is there as it is. +func (p *localPlan) dataDir(ctx context.Context) error { + var owned []Owned + dirs := []string{DataDir} + if p.h.OS == "darwin" { + dirs = append(dirs, macLogDir) + } + for _, dir := range dirs { + made, err := p.makeDir(ctx, dir) + if err != nil { + return &Problem{Lines: []string{fmt.Sprintf("Could not create %s: %v.", dir, err)}} + } + if made { + owned = append(owned, Owned{Path: dir, Kind: KindData}) + } + } + + return RecordOwned(ctx, p.h, p.local, owned...) +} + +func (p *localPlan) makeDir(ctx context.Context, dir string) (bool, error) { + _, err := os.Lstat(rooted(p.h, dir)) + if err == nil { + return false, nil + } + if !errors.Is(err, fs.ErrNotExist) { + return false, err + } + // The daemon's own MkdirAll makes 0750 root on Linux; on a Mac the daemon runs as the person, who owns its root. + args := []string{"-d", "-o", "0", "-g", "0", "-m", "0750", rooted(p.h, dir)} + if p.h.OS == "darwin" { + args = []string{"-d", "-o", p.user, "-m", "0755", rooted(p.h, dir)} + } + if _, err := privileged(ctx, p.h, "install", args...); err != nil { + return false, err + } + + return true, nil +} + +// service writes the unit or the LaunchDaemon, and has the service manager check it before it is enabled. +func (p *localPlan) service(ctx context.Context) (err error) { + dir, err := os.MkdirTemp("", "shard-setup-") + if err != nil { + return fmt.Errorf("create a staging directory: %w", err) + } + defer func() { + if rmErr := os.RemoveAll(dir); rmErr != nil { + err = errors.Join(err, fmt.Errorf("remove the staging directory: %w", rmErr)) + } + }() + + if p.h.OS == "darwin" { + return p.launchdService(ctx, dir) + } + + return p.systemdService(ctx, dir) +} + +func (p *localPlan) systemdService(ctx context.Context, dir string) error { + unit := filepath.Join(dir, "shard.service") + if err := os.WriteFile(unit, []byte(systemdUnitText(p.local.Provider)), 0o600); err != nil { + return err + } + if _, err := run(ctx, p.h, "systemd-analyze", "verify", unit); err != nil { + return &Problem{Lines: []string{fmt.Sprintf("systemd refused the service definition: %v.", err)}} + } + if err := installFile(ctx, p.h, unit, systemdUnit, "0644"); err != nil { + return &Problem{Lines: []string{fmt.Sprintf("Could not install %s: %v.", systemdUnit, err)}} + } + for _, args := range [][]string{{"daemon-reload"}, {"enable", "shard.service"}} { + if _, err := privileged(ctx, p.h, "systemctl", args...); err != nil { + return &Problem{Lines: []string{fmt.Sprintf("Could not enable the service: %v.", err)}} + } + } + + return RecordOwned(ctx, p.h, p.local, Owned{Path: systemdUnit, Kind: KindService}) +} + +func (p *localPlan) launchdService(ctx context.Context, dir string) error { + plist := filepath.Join(dir, "shard.daemon.plist") + if err := os.WriteFile(plist, []byte(launchdPlistText(p.user)), 0o600); err != nil { + return err + } + if _, err := run(ctx, p.h, "plutil", "-lint", plist); err != nil { + return &Problem{Lines: []string{fmt.Sprintf("launchd refused the service definition: %v.", err)}} + } + rotate := filepath.Join(dir, "shard.conf") + if err := os.WriteFile(rotate, []byte(newsyslogText(p.user)), 0o600); err != nil { + return err + } + for _, f := range []struct{ src, dst string }{{plist, launchdPlist}, {rotate, newsyslog}} { + if err := installFile(ctx, p.h, f.src, f.dst, "0644"); err != nil { + return &Problem{Lines: []string{fmt.Sprintf("Could not install %s: %v.", f.dst, err)}} + } + } + + return RecordOwned(ctx, p.h, p.local, Owned{Path: launchdPlist, Kind: KindService}, Owned{Path: newsyslog, Kind: KindConfig}) +} + +func (p *localPlan) start(ctx context.Context) error { + if p.h.OS != "darwin" { + if _, err := privileged(ctx, p.h, "systemctl", "start", "shard.service"); err != nil { + return &Problem{Lines: []string{fmt.Sprintf("Could not start the daemon: %v.", err)}} + } + + return nil + } + + // A retry finds the job already loaded, and bootstrap refuses a loaded one. + if _, err := privileged(ctx, p.h, "launchctl", "print", launchdLabel); err != nil { + if _, err := privileged(ctx, p.h, "launchctl", "bootstrap", "system", rooted(p.h, launchdPlist)); err != nil { + return &Problem{Lines: []string{fmt.Sprintf("Could not start the daemon: %v.", err)}} + } + + return nil + } + if _, err := privileged(ctx, p.h, "launchctl", "kickstart", launchdLabel); err != nil { + return &Problem{Lines: []string{fmt.Sprintf("Could not start the daemon: %v.", err)}} + } + + return nil +} + +// verify asks the daemon for its status until it answers, through sudo on Linux where the socket is root's; --remote "" keeps any remote out of it. +func (p *localPlan) verify(ctx context.Context) error { + ask := privileged + if p.h.OS == "darwin" { + ask = run + } + deadline := time.Now().Add(verifyWait) + for { + out, err := ask(ctx, p.h, shardBinary, "--remote", "", "daemon", "status") + if err == nil { + return nil + } + if time.Now().After(deadline) { + return &Problem{Lines: []string{ + fmt.Sprintf("The daemon is not ready after %s: %s.", verifyWait, notReady(out, err)), + p.logHint(), + }} + } + select { + case <-ctx.Done(): + return ctx.Err() + case <-time.After(verifyPoll): + } + } +} + +// notReady is the reason daemon status printed, without the command line behind it, else the error whole. +func notReady(out []byte, err error) string { + if reason := strings.TrimPrefix(strings.TrimPrefix(outputTail(out), ": "), "shard: "); reason != "" { + return reason + } + + return err.Error() +} + +func (p *localPlan) logHint() string { + if p.h.OS == "darwin" { + return "Read its log in " + macLogDir + "/daemon.log." + } + + return "Read its log with: sudo journalctl -u shard.service" +} + +// systemdUnitText is packaging/systemd/shard.service with the provider named, so the daemon never probes for one. +func systemdUnitText(provider string) string { + return strings.Replace(systemdUnitTemplate, systemdExecStart, systemdExecStart+" --provider "+provider, 1) +} + +const systemdExecStart = "ExecStart=" + shardBinary + " daemon" + +// launchdPlistText is packaging/launchd/shard.daemon.plist with the person's name put in. +func launchdPlistText(user string) string { + return strings.ReplaceAll(launchdPlistTemplate, "__USER__", user) +} + +func newsyslogText(user string) string { + return strings.ReplaceAll(newsyslogTemplate, "__USER__", user) +} + +const systemdUnitTemplate = `# shard daemon is a resident root process, installed on purpose: no one-shot verb ever spawns it. +[Unit] +Description=shard daemon, the background work beside the CLI +After=network-online.target +Wants=network-online.target + +[Service] +Type=simple +ExecStart=/usr/local/bin/shard daemon +# A unix socket needs write to connect, so it must never sit at 0755 between listen and chmod. +UMask=0077 +# runsc sits in its own session, and a sandbox outlives the daemon: a restart or a stop kills the daemon alone. +KillMode=process +Restart=on-failure +RestartSec=1 + +[Install] +WantedBy=multi-user.target +` + +const launchdPlistTemplate = `<?xml version="1.0" encoding="UTF-8"?> +<!DOCTYPE plist PUBLIC "-//Apple//DTD PLIST 1.0//EN" "http://www.apple.com/DTDs/PropertyList-1.0.dtd"> +<!-- shard daemon is a resident process, installed on purpose: no one-shot verb ever spawns it. The LaunchDaemon mirror of packaging/systemd/shard.service. --> +<plist version="1.0"> +<dict> + <key>Label</key> + <string>shard.daemon</string> + <key>ProgramArguments</key> + <array> + <string>/usr/local/bin/shard</string> + <string>daemon</string> + <!-- launchd opens the log once, so the daemon opens it itself and reopens it on the SIGHUP newsyslog sends. --> + <string>--log</string> + <string>/var/log/shard/daemon.log</string> + </array> + <!-- The root under /var/lib/shard belongs to this user, so nothing runs as root; docs/mac.md puts the name in. --> + <key>UserName</key> + <string>__USER__</string> + <key>RunAtLoad</key> + <true/> + <!-- Restart=on-failure: a clean exit stays down, a crash comes back after ThrottleInterval, as RestartSec=1. --> + <key>KeepAlive</key> + <dict> + <key>SuccessfulExit</key> + <false/> + </dict> + <key>ThrottleInterval</key> + <integer>1</integer> + <!-- KillMode=process: a sandbox outlives the daemon, so a stop or a restart ends the daemon and leaves every VM shim to be adopted. --> + <key>AbandonProcessGroup</key> + <true/> + <!-- UMask=0077, as a decimal: a unix socket needs write to connect, so it must never sit at 0755 between listen and chmod. --> + <key>Umask</key> + <integer>63</integer> + <!-- The same file, for what the daemon prints before --log takes over, and a refusal to start. --> + <key>StandardOutPath</key> + <string>/var/log/shard/daemon.log</string> + <key>StandardErrorPath</key> + <string>/var/log/shard/daemon.log</string> +</dict> +</plist> +` + +const newsyslogTemplate = `# logfilename [owner:group] mode count size(KB) when flags pid_file sig_num +# macOS newsyslog keeps .0 to .count, so a count of 6 keeps seven old files. +/var/log/shard/daemon.log __USER__:staff 600 6 10240 * - /var/lib/shard/daemon.pid 1 +` diff --git a/services/setup/existing.go b/services/setup/existing.go new file mode 100644 index 000000000..12908fd20 --- /dev/null +++ b/services/setup/existing.go @@ -0,0 +1,602 @@ +package setup + +import ( + "bytes" + "context" + "errors" + "fmt" + "io/fs" + "os" + "path/filepath" + "slices" + "strings" + + "github.com/presmihaylov/shard/pkg/term" + "github.com/presmihaylov/shard/services/sandboxstate" +) + +const serviceName = "shard" + +// manualPaths are where docs/daemon.md and docs/mac.md put an install by hand. +var manualPaths = []string{shardBinary, initBinary, systemdUnit, "/etc/systemd/system/shard-serve.service", launchdPlist, newsyslog} + +// ServiceState is the Service line of the summary. +type ServiceState string + +const ( + ServiceActive ServiceState = "Active" + ServiceInactive ServiceState = "Inactive" + ServiceNone ServiceState = "Not set up" +) + +// Installation is what Detect found of an earlier local install. +type Installation struct { + // Manifest is set when setup made the install; Manual lists what an install by hand left instead. + Manifest *Manifest + Manual []string + Service ServiceState +} + +// Detect reads the host without changing it and without administrator access. +func Detect(ctx context.Context, h Host) (Installation, bool, error) { + m, ok, err := LoadManifest(h) + if err != nil { + return Installation{}, false, err + } + if ok { + state, err := serviceState(ctx, h, m) + if err != nil { + return Installation{}, false, err + } + + return Installation{Manifest: &m, Service: state}, true, nil + } + + var found []string + for _, p := range manualPaths { + _, err := os.Lstat(filepath.Join(h.Root, p)) + if errors.Is(err, fs.ErrNotExist) { + continue + } + if err != nil { + return Installation{}, false, fmt.Errorf("check %s: %w", p, err) + } + found = append(found, p) + } + + return Installation{Manual: found}, len(found) > 0, nil +} + +// existing is the §11 menu over an installation Detect found. +func (s *Setup) existing(ctx context.Context, inst Installation) error { + if inst.Manifest == nil { + return s.manual(inst) + } + + m := *inst.Manifest + if err := s.UI.Print("Shard is already installed", "", "Version: "+m.Version, "Provider: "+providerTitle(m.Provider), "Service: "+string(inst.Service), ""); err != nil { + return err + } + choice, err := s.UI.Select(ctx, AskExisting, "What would you like to do?", []term.Option{ + {Name: "repair", Label: "Check or repair the installation", Default: true}, + {Name: "upgrade", Label: "Upgrade Shard"}, + {Name: "uninstall", Label: "Uninstall Shard"}, + {Name: "exit", Label: "Exit"}, + }) + if err != nil { + return err + } + + switch choice { + case 0: + return s.switched(ctx, func(ctx context.Context) error { return s.repair(ctx, m, inst.Service) }) + case 1: + return s.switched(ctx, func(ctx context.Context) error { return s.upgrade(ctx, m, inst.Service) }) + case 2: + return s.uninstall(ctx, m) + } + + return nil +} + +// repair re-runs the steps of the recorded choices, so the provider and the startup setting stay what they were. +func (s *Setup) repair(ctx context.Context, m Manifest, service ServiceState) error { + var problems []string + for _, f := range m.Files { + _, err := os.Lstat(filepath.Join(s.Host.Root, f.Path)) + if errors.Is(err, fs.ErrNotExist) { + problems = append(problems, f.Path+" is missing.") + continue + } + if err != nil { + return fmt.Errorf("check %s: %w", f.Path, err) + } + } + if m.StartAtBoot && service != ServiceActive { + problems = append(problems, "The background service is not running.") + } + + if len(problems) == 0 { + return s.UI.Print("✓ No problems found", "", "Shard "+m.Version+" with "+providerTitle(m.Provider)+" is installed correctly.") + } + + steps, err := s.localSteps(ctx, Local{Provider: m.Provider, StartAtBoot: m.StartAtBoot}) + if err != nil { + return err + } + lines := []string{"Setup found these problems:", ""} + for _, p := range problems { + lines = append(lines, " "+p) + } + lines = append(lines, "", "Setup will run these steps again:") + for _, st := range steps { + lines = append(lines, " "+st.Title) + } + lines = append(lines, "", "Your settings and sandbox data stay in place.", "") + if err := s.UI.Print(lines...); err != nil { + return err + } + if err := s.confirm(ctx, true); err != nil { + return err + } + if err := s.admin(ctx); err != nil { + return err + } + + return s.apply(ctx, "Repairing Shard", steps) +} + +// replacement is one installed binary and the release file that replaces it. +type replacement struct { + path, asset string + // user marks the CLI the user runs, which the user owns and replaces without administrator access. + user bool + tmp string +} + +// upgrade fetches and verifies every binary before it asks, and replaces none until all of them passed. +func (s *Setup) upgrade(ctx context.Context, m Manifest, service ServiceState) (err error) { + h := s.Host + rel, err := LatestRelease(ctx, h) + if err != nil { + return err + } + if err := s.UI.Print("Latest release: " + rel.Tag); err != nil { + return err + } + latest, _ := stableVersion(rel.Tag) + if installed, ok := stableVersion(m.Version); ok && !newer(latest, installed) { + return s.UI.Print("", "Shard "+m.Version+" is up to date.") + } + + targets, err := upgradeTargets(h, m) + if err != nil { + return err + } + dir, err := os.MkdirTemp("", "shard-upgrade-*") + if err != nil { + return fmt.Errorf("upgrade Shard: %w", err) + } + defer func() { err = errors.Join(err, os.RemoveAll(dir)) }() + + fetch := Step{Title: "Download and verify Shard " + rel.Tag, Do: func(ctx context.Context) error { + fetched := map[string]bool{} + for i := range targets { + t := &targets[i] + t.tmp = filepath.Join(dir, t.asset) + if fetched[t.asset] { + continue + } + if err := rel.Fetch(ctx, h, t.asset, t.tmp, 0o755); err != nil { + return err + } + fetched[t.asset] = true + } + + return verifyVersion(ctx, h, targets, rel.Tag) + }} + if err := s.apply(ctx, "Preparing the upgrade", []Step{fetch}); err != nil { + return err + } + + lines := []string{"", "Setup will replace:"} + for _, t := range targets { + lines = append(lines, " "+t.path) + } + lines = append(lines, "", "The provider stays "+providerTitle(m.Provider)+".") + switch service { + case ServiceActive: + lines = append(lines, "The background service restarts to run "+rel.Tag+".", + "Your sandboxes keep running while the daemon restarts.", + "Open `shard exec` sessions disconnect. Their commands keep running.") + case ServiceInactive: + lines = append(lines, "The background service is not running. Setup does not start it.") + case ServiceNone: + lines = append(lines, "You start the daemon yourself. Restart `shard daemon` to run "+rel.Tag+".") + } + if err := s.UI.Print(append(lines, "")...); err != nil { + return err + } + if err := s.confirm(ctx, true); err != nil { + return err + } + if err := s.admin(ctx); err != nil { + return err + } + + steps := []Step{ + {Title: "Replace the Shard binaries", Do: func(ctx context.Context) error { return replaceAll(ctx, h, targets) }}, + {Title: "Record Shard " + rel.Tag, Do: func(ctx context.Context) error { + m.Version = rel.Tag + return saveManifest(ctx, h, m) + }}, + } + if service == ServiceActive { + steps = append(steps, Step{Title: "Restart the daemon", Do: func(ctx context.Context) error { return restartService(ctx, h, m) }}) + } + + return s.apply(ctx, "Upgrading Shard", steps) +} + +// upgradeTargets is every binary the manifest owns, and the CLI that runs setup when the manifest does not own it. +func upgradeTargets(h Host, m Manifest) ([]replacement, error) { + var targets []replacement + for _, f := range m.Files { + if f.Kind != KindBinary { + continue + } + asset, err := releaseAsset(h, filepath.Base(f.Path)) + if err != nil { + return nil, err + } + targets = append(targets, replacement{path: f.Path, asset: asset}) + } + + owned := slices.ContainsFunc(targets, func(r replacement) bool { return filepath.Join(h.Root, r.path) == h.Executable }) + if !owned { + targets = append(targets, replacement{path: h.Executable, asset: "shard-" + h.OS + "-" + h.Arch, user: true}) + } + + return targets, nil +} + +func releaseAsset(h Host, name string) (string, error) { + switch name { + case "shard", "shard-init": + return name + "-" + h.OS + "-" + h.Arch, nil + } + + return "", fmt.Errorf("upgrade Shard: no release file replaces %s", name) +} + +// verifyVersion runs each new CLI before anything is replaced; shard-init is a guest PID 1, so its checksum is its proof. +func verifyVersion(ctx context.Context, h Host, targets []replacement, tag string) error { + for _, t := range targets { + if strings.HasPrefix(t.asset, "shard-init-") { + continue + } + out, err := run(ctx, h, t.tmp, "--version") + if err != nil { + return fmt.Errorf("run the new %s: %w", t.asset, err) + } + if got := strings.TrimSpace(string(out)); got != "client "+tag { + return fmt.Errorf("the new %s reports %q, want %q", t.asset, got, "client "+tag) + } + } + + return nil +} + +// replaceAll stages every file beside its target first, so one failed copy leaves no mix of versions. +func replaceAll(ctx context.Context, h Host, targets []replacement) error { + for _, t := range targets { + dst := filepath.Join(h.Root, t.path) + if t.user { + dst = t.path + } + if err := stage(ctx, h, t, dst+".new"); err != nil { + return fmt.Errorf("stage %s: %w", t.path, err) + } + } + + for _, t := range targets { + dst := filepath.Join(h.Root, t.path) + if t.user { + if err := os.Rename(t.path+".new", t.path); err != nil { + return fmt.Errorf("replace %s: %w", t.path, err) + } + continue + } + if _, err := privileged(ctx, h, "mv", "-f", dst+".new", dst); err != nil { + return fmt.Errorf("replace %s: %w", t.path, err) + } + } + + return nil +} + +func stage(ctx context.Context, h Host, t replacement, next string) error { + if t.user { + data, err := os.ReadFile(t.tmp) + if err != nil { + return err + } + + return os.WriteFile(next, data, 0o755) //nolint:gosec // the CLI the user runs must stay executable + } + + _, err := privileged(ctx, h, "install", "-m", "0755", t.tmp, next) + + return err +} + +func restartService(ctx context.Context, h Host, m Manifest) error { + cmd := []string{"systemctl", "restart", serviceName} + if h.OS == "darwin" { + cmd = []string{"launchctl", "kickstart", "-k", launchdLabel} + } + if _, err := privileged(ctx, h, cmd[0], cmd[1:]...); err != nil { + return fmt.Errorf("restart the daemon: %w", err) + } + + state, err := serviceState(ctx, h, m) + if err != nil { + return err + } + if state != ServiceActive { + return fmt.Errorf("the daemon is %s after the restart", strings.ToLower(string(state))) + } + + return nil +} + +// uninstall refuses while a sandbox is left, and keeps saved data and shared tools. +func (s *Setup) uninstall(ctx context.Context, m Manifest) error { + h := s.Host + if err := s.admin(ctx); err != nil { + return err + } + n, err := countSandboxes(ctx, h) + if err != nil { + return err + } + if n > 0 { + left := fmt.Sprintf("%d %s", n, plural(n, "sandbox", "sandboxes")) + return errors.Join(s.UI.Print("Shard has "+left+" on this machine.", + "Remove them before you uninstall Shard:", "", + " List sandboxes:", " shard list --all", "", + " Remove a sandbox:", " shard remove --force <name>"), + fmt.Errorf("uninstall stopped: %s left", left)) + } + + if err := s.UI.Print("Uninstall Shard?", "", + "This will stop and remove the background service", + "and remove files installed by Shard setup.", "", + "Your saved data will remain.", + "Shared tools will remain.", ""); err != nil { + return err + } + if err := s.confirm(ctx, false); err != nil { + return err + } + + steps := []Step{ + {Title: "Stop and remove the background service", Do: func(ctx context.Context) error { return stopService(ctx, h, m) }}, + {Title: "Remove files installed by Shard setup", Do: func(ctx context.Context) error { return removeOwned(ctx, h, m) }}, + } + if err := s.apply(ctx, "Uninstalling Shard", steps); err != nil { + return err + } + + return s.UI.Print(uninstalled(h, m)...) +} + +// countSandboxes reads the records with administrator access, since the data root is not the user's on Linux. +func countSandboxes(ctx context.Context, h Host) (int, error) { + root := filepath.Join(h.Root, DataDir) + _, err := os.Lstat(root) + if errors.Is(err, fs.ErrNotExist) { + return 0, nil + } + if err != nil { + return 0, fmt.Errorf("check %s: %w", DataDir, err) + } + + out, err := privileged(ctx, h, "find", root, "-mindepth", "3", "-maxdepth", "3", "-name", "sandbox.json") + if err != nil { + return 0, fmt.Errorf("list the sandboxes in %s: %w", DataDir, err) + } + + n := 0 + for line := range strings.SplitSeq(strings.TrimSpace(string(out)), "\n") { + dir := filepath.Dir(line) + if filepath.Base(filepath.Dir(dir)) == "sandboxes" && sandboxstate.ValidID(filepath.Base(dir)) == nil { + n++ + } + } + + return n, nil +} + +// stopService checks before each stop, so a retry after a half-done uninstall stops nothing twice. +func stopService(ctx context.Context, h Host, m Manifest) error { + var units []string + for _, f := range m.Files { + if f.Kind != KindService { + continue + } + _, err := os.Lstat(filepath.Join(h.Root, f.Path)) + if errors.Is(err, fs.ErrNotExist) { + continue + } + if err != nil { + return fmt.Errorf("check %s: %w", f.Path, err) + } + units = append(units, f.Path) + } + if len(units) == 0 { + return nil + } + + if h.OS == "darwin" { + out, err := h.Run(ctx, "launchctl", "print", launchdLabel) + if err != nil && bytes.Contains(out, []byte("Could not find service")) { + return nil + } + if err != nil { + return fmt.Errorf("check the background service: %w%s", err, outputTail(out)) + } + if _, err := privileged(ctx, h, "launchctl", "bootout", launchdLabel); err != nil { + return fmt.Errorf("stop the background service: %w", err) + } + + return nil + } + + for _, u := range units { + name := strings.TrimSuffix(filepath.Base(u), ".service") + if _, err := privileged(ctx, h, "systemctl", "disable", "--now", name); err != nil { + return fmt.Errorf("stop %s: %w", name, err) + } + } + + return nil +} + +// removeOwned removes shard's own files newest first, and the manifest last, so a failed run can run again. +func removeOwned(ctx context.Context, h Host, m Manifest) error { + service := false + for _, f := range slices.Backward(m.Files) { + if f.Kind != KindBinary && f.Kind != KindService && f.Kind != KindConfig { + continue + } + if _, err := privileged(ctx, h, "rm", "-f", filepath.Join(h.Root, f.Path)); err != nil { + return fmt.Errorf("remove %s: %w", f.Path, err) + } + service = service || f.Kind == KindService + } + if service && h.OS == "linux" { + if _, err := privileged(ctx, h, "systemctl", "daemon-reload"); err != nil { + return fmt.Errorf("reload systemd: %w", err) + } + } + + return removeManifest(ctx, h) +} + +// uninstalled names what stays and how to remove it, since uninstall removes no shared tool and no data. +func uninstalled(h Host, m Manifest) []string { + lines := []string{"", "Shard was uninstalled.", "", "Your saved data remains in " + DataDir + "."} + + var tools []string + for _, f := range m.Files { + if f.Kind != KindTool { + continue + } + if f.Package != "" { + tools = append(tools, " "+f.Package, " Remove it with: sudo apt-get remove "+f.Package) + continue + } + tools = append(tools, " "+f.Path, " Remove it with: sudo rm "+f.Path) + } + if len(tools) > 0 { + lines = append(lines, "", "These shared tools remain:") + lines = append(lines, tools...) + } + + return append(lines, "", "The shard command remains at "+h.Executable+".", "Remove it with: rm "+h.Executable) +} + +// manual reports an install setup did not make, and changes nothing in it. +func (s *Setup) manual(inst Installation) error { + lines := []string{"Manual installation detected.", "", "Found:"} + var bins, units, others []string + for _, p := range inst.Manual { + lines = append(lines, " "+p) + switch { + case strings.HasSuffix(p, ".service"): + units = append(units, p) + case strings.HasPrefix(p, "/usr/local/bin/"): + bins = append(bins, p) + default: + others = append(others, p) + } + } + lines = append(lines, "", "Setup did not install these files, so it does not change or remove them.", "", "Inspect the installation:", " shard version") + if s.Host.OS == "darwin" { + lines = append(lines, " sudo launchctl print "+launchdLabel) + } + if s.Host.OS == "linux" { + lines = append(lines, " systemctl status shard") + } + + lines = append(lines, "", "To remove it, remove every sandbox first (shard list --all), then run:") + for _, u := range units { + lines = append(lines, " sudo systemctl disable --now "+strings.TrimSuffix(filepath.Base(u), ".service")) + } + if slices.Contains(others, "/Library/LaunchDaemons/shard.daemon.plist") { + lines = append(lines, " sudo launchctl bootout "+launchdLabel) + } + if files := slices.Concat(units, others, bins); len(files) > 0 { + lines = append(lines, " sudo rm "+strings.Join(files, " ")) + } + if len(units) > 0 { + lines = append(lines, " sudo systemctl daemon-reload") + } + + return s.UI.Print(append(lines, "", "Your saved data in "+DataDir+" is not part of this.")...) +} + +func serviceState(ctx context.Context, h Host, m Manifest) (ServiceState, error) { + if !m.StartAtBoot { + return ServiceNone, nil + } + + if h.OS == "darwin" { + out, err := h.Run(ctx, "launchctl", "print", launchdLabel) + if err != nil && bytes.Contains(out, []byte("Could not find service")) { + return ServiceInactive, nil + } + if err != nil { + return "", fmt.Errorf("check the background service: %w%s", err, outputTail(out)) + } + if bytes.Contains(out, []byte("state = running")) { + return ServiceActive, nil + } + + return ServiceInactive, nil + } + + // is-active exits non-zero for every state but active, and still prints the state. + out, err := h.Run(ctx, "systemctl", "is-active", serviceName) + state := strings.TrimSpace(string(out)) + if state == "active" { + return ServiceActive, nil + } + if err != nil && state == "" { + return "", fmt.Errorf("check the background service: %w", err) + } + + return ServiceInactive, nil +} + +// confirm asks before a change; a no is ErrDeclined, so the run ends having changed nothing. +func (s *Setup) confirm(ctx context.Context, yes bool) error { + ok, err := s.UI.Confirm(ctx, AskConfirm, "Continue?", yes) + if err != nil { + return err + } + if !ok { + return ErrDeclined + } + + return nil +} + +func plural(n int, one, many string) string { + if n == 1 { + return one + } + + return many +} diff --git a/services/setup/existing_test.go b/services/setup/existing_test.go new file mode 100644 index 000000000..03a6d458d --- /dev/null +++ b/services/setup/existing_test.go @@ -0,0 +1,548 @@ +package setup + +import ( + "context" + "errors" + "os" + "os/exec" + "path/filepath" + "slices" + "strings" + "testing" + + "github.com/presmihaylov/shard/services/client" +) + +// fakeHost runs file commands for real under a temp root and answers the service managers from canned output. +type fakeHost struct { + t *testing.T + root string + calls []string + isActive string + launchd string + // launchdErr makes launchctl print fail with launchd as its output. + launchdErr bool +} + +func newFakeHost(t *testing.T) *fakeHost { + t.Helper() + return &fakeHost{t: t, root: t.TempDir(), isActive: "active"} +} + +// host runs as root, so privileged runs each command as it is. +func (f *fakeHost) host(rs *releaseServer) Host { + h := Host{Root: f.root, OS: "linux", Arch: "amd64", Executable: filepath.Join(f.root, "/home/u/.local/bin/shard"), Version: "v0.1.0", Run: f.run} + if rs != nil { + h.Releases, h.HTTP = rs.URL+"/releases", rs.Client() + } + + return h +} + +func (f *fakeHost) run(ctx context.Context, name string, args ...string) ([]byte, error) { + f.calls = append(f.calls, strings.Join(append([]string{name}, args...), " ")) + + switch name { + case "systemctl": + if args[0] != "is-active" { + return nil, nil + } + if f.isActive == "active" { + return []byte("active\n"), nil + } + return []byte(f.isActive + "\n"), errors.New("exit status 3") + case "launchctl": + if args[0] == "print" && f.launchdErr { + return []byte(f.launchd), errors.New("exit status 113") + } + return []byte(f.launchd), nil + case "mkdir", "install", "mv", "rm", "rmdir", "find": + return exec.CommandContext(ctx, name, args...).CombinedOutput() + } + if filepath.IsAbs(name) { + return exec.CommandContext(ctx, name, args...).CombinedOutput() + } + f.t.Errorf("unexpected command %s %v", name, args) + + return nil, errors.New("unexpected command") +} + +func (f *fakeHost) called(prefix string) bool { + return slices.ContainsFunc(f.calls, func(c string) bool { return strings.HasPrefix(c, prefix) }) +} + +func (f *fakeHost) write(t *testing.T, path, body string) { + t.Helper() + full := filepath.Join(f.root, path) + if err := os.MkdirAll(filepath.Dir(full), 0o755); err != nil { + t.Fatalf("mkdir %s: %v", path, err) + } + if err := os.WriteFile(full, []byte(body), 0o755); err != nil { + t.Fatalf("write %s: %v", path, err) + } +} + +func (f *fakeHost) read(t *testing.T, path string) (string, bool) { + t.Helper() + data, err := os.ReadFile(filepath.Join(f.root, path)) + if errors.Is(err, os.ErrNotExist) { + return "", false + } + if err != nil { + t.Fatalf("read %s: %v", path, err) + } + + return string(data), true +} + +// installed writes the files of a setup install, the manifest that names them, and the shard the user runs. +func (f *fakeHost) installed(t *testing.T, m Manifest) { + t.Helper() + f.write(t, "/home/u/.local/bin/shard", "old cli") + for _, o := range m.Files { + if o.Kind != KindData { + f.write(t, o.Path, "old "+filepath.Base(o.Path)) + } + } + if err := saveManifest(t.Context(), f.host(nil), m); err != nil { + t.Fatalf("save manifest: %v", err) + } + f.calls = nil +} + +func linuxInstall(version string) Manifest { + return Manifest{Version: version, Provider: "gvisor", StartAtBoot: true, Files: []Owned{ + {Path: "/usr/sbin/nft", Kind: KindTool, Package: "nftables"}, + {Path: "/usr/local/bin/shard", Kind: KindBinary}, + {Path: "/usr/local/bin/shard-init", Kind: KindBinary}, + {Path: "/etc/systemd/system/shard.service", Kind: KindService}, + {Path: "/etc/shard/serve.env", Kind: KindConfig}, + }} +} + +// watchUI records the default of each Confirm and calls onPrint before it keeps a printed line. +type watchUI struct { + *fakeUI + defaults []bool + onPrint func(lines []string) +} + +func (u *watchUI) Confirm(ctx context.Context, q Question, text string, yes bool) (bool, error) { + u.defaults = append(u.defaults, yes) + return u.fakeUI.Confirm(ctx, q, text, yes) +} + +func (u *watchUI) Print(lines ...string) error { + if u.onPrint != nil { + u.onPrint(lines) + } + return u.fakeUI.Print(lines...) +} + +func confirming(yes bool) *fakeUI { + return &fakeUI{confirms: map[Question]bool{AskConfirm: yes}} +} + +func said(t *testing.T, ui *fakeUI, lines ...string) { + t.Helper() + for _, line := range lines { + if !slices.ContainsFunc(ui.printed, func(p string) bool { return strings.Contains(p, line) }) { + t.Errorf("output %q lacks %q", ui.printed, line) + } + } +} + +func TestDetect(t *testing.T) { + t.Run("nothing installed", func(t *testing.T) { + f := newFakeHost(t) + f.write(t, "/var/lib/shard/images/index.json", "{}") + + _, ok, err := Detect(t.Context(), f.host(nil)) + if err != nil || ok { + t.Fatalf("Detect = %v, %v; the data dir alone is no installation", ok, err) + } + }) + + t.Run("setup install", func(t *testing.T) { + f := newFakeHost(t) + f.installed(t, linuxInstall("v0.1.0")) + f.isActive = "failed" + + inst, ok, err := Detect(t.Context(), f.host(nil)) + if err != nil || !ok { + t.Fatalf("Detect = %v, %v", ok, err) + } + if inst.Manifest == nil || inst.Manifest.Version != "v0.1.0" || inst.Service != ServiceInactive { + t.Fatalf("Detect = %+v, want the manifest and an inactive service", inst) + } + }) + + t.Run("manual startup asks no service manager", func(t *testing.T) { + f := newFakeHost(t) + m := linuxInstall("v0.1.0") + m.StartAtBoot = false + f.installed(t, m) + + inst, _, err := Detect(t.Context(), f.host(nil)) + if err != nil || inst.Service != ServiceNone || len(f.calls) != 0 { + t.Fatalf("Detect = %+v, %v, calls %v", inst, err, f.calls) + } + }) + + t.Run("manual install", func(t *testing.T) { + f := newFakeHost(t) + f.write(t, "/usr/local/bin/shard", "bin") + f.write(t, "/etc/systemd/system/shard.service", "unit") + + inst, ok, err := Detect(t.Context(), f.host(nil)) + if err != nil || !ok || inst.Manifest != nil { + t.Fatalf("Detect = %+v, %v, %v", inst, ok, err) + } + want := []string{"/usr/local/bin/shard", "/etc/systemd/system/shard.service"} + if !slices.Equal(inst.Manual, want) { + t.Fatalf("Manual = %v, want %v", inst.Manual, want) + } + }) +} + +func TestServiceStateOnMac(t *testing.T) { + cases := map[string]struct { + out string + fails bool + want ServiceState + wantErr bool + }{ + "running": {out: "system/shard.daemon = {\n\tstate = running\n}", want: ServiceActive}, + "loaded": {out: "system/shard.daemon = {\n\tstate = not running\n}", want: ServiceInactive}, + "not found": {out: "Could not find service \"shard.daemon\" in domain for system", fails: true, want: ServiceInactive}, + "other": {out: "Operation not permitted", fails: true, wantErr: true}, + } + for name, tc := range cases { + t.Run(name, func(t *testing.T) { + f := newFakeHost(t) + f.launchd, f.launchdErr = tc.out, tc.fails + h := f.host(nil) + h.OS = "darwin" + + got, err := serviceState(t.Context(), h, Manifest{StartAtBoot: true}) + if (err != nil) != tc.wantErr || got != tc.want { + t.Fatalf("serviceState = %q, %v; want %q, error %v", got, err, tc.want, tc.wantErr) + } + }) + } +} + +func TestExistingShowsTheSummaryAndExits(t *testing.T) { + f := newFakeHost(t) + m := linuxInstall("v0.1.0") + f.installed(t, m) + ui := &fakeUI{selects: map[Question]string{AskExisting: "exit"}} + + err := (&Setup{Host: f.host(nil), UI: ui}).existing(t.Context(), Installation{Manifest: &m, Service: ServiceActive}) + if err != nil { + t.Fatalf("existing: %v", err) + } + want := []string{"Shard is already installed", "", "Version: v0.1.0", "Provider: " + providerTitle("gvisor"), "Service: Active", ""} + if !slices.Equal(ui.printed, want) { + t.Fatalf("summary = %q, want %q", ui.printed, want) + } + var names []string + for _, o := range ui.options[AskExisting] { + names = append(names, o.Name) + } + if !slices.Equal(names, []string{"repair", "upgrade", "uninstall", "exit"}) || !ui.options[AskExisting][0].Default { + t.Fatalf("the menu offers %v, want repair (the default), upgrade, uninstall, exit", names) + } + if len(f.calls) != 0 { + t.Fatalf("Exit ran %v", f.calls) + } +} + +func TestManualInstallChangesNothing(t *testing.T) { + f := newFakeHost(t) + f.write(t, "/usr/local/bin/shard", "bin") + f.write(t, "/etc/systemd/system/shard.service", "unit") + inst, _, err := Detect(t.Context(), f.host(nil)) + if err != nil { + t.Fatalf("Detect: %v", err) + } + ui := &fakeUI{} + + if err := (&Setup{Host: f.host(nil), UI: ui}).existing(t.Context(), inst); err != nil { + t.Fatalf("existing: %v", err) + } + said(t, ui, "Manual installation detected.", "sudo systemctl disable --now shard", "sudo rm /etc/systemd/system/shard.service /usr/local/bin/shard") + if len(f.calls) != 0 || len(ui.asked) != 0 { + t.Fatalf("a manual install ran %v and asked %v", f.calls, ui.asked) + } + if _, ok := f.read(t, "/usr/local/bin/shard"); !ok { + t.Fatal("a manual install lost its binary") + } +} + +func TestRepairFindsNothingToDo(t *testing.T) { + f := newFakeHost(t) + m := linuxInstall("v0.1.0") + f.installed(t, m) + ui := &fakeUI{} + + if err := (&Setup{Host: f.host(nil), UI: ui}).repair(t.Context(), m, ServiceActive); err != nil { + t.Fatalf("repair: %v", err) + } + said(t, ui, "No problems found") + if len(ui.asked) != 0 { + t.Fatalf("repair asked %v", ui.asked) + } +} + +func TestRepairShowsTheProblemsBeforeItAsks(t *testing.T) { + f := newFakeHost(t) + m := linuxInstall("v0.1.0") + f.installed(t, m) + if err := os.Remove(filepath.Join(f.root, "/usr/local/bin/shard-init")); err != nil { + t.Fatalf("remove: %v", err) + } + ui := confirming(false) + + err := (&Setup{Host: f.host(nil), UI: ui}).repair(t.Context(), m, ServiceInactive) + if !errors.Is(err, ErrDeclined) { + t.Fatalf("repair = %v, want ErrDeclined", err) + } + said(t, ui, "/usr/local/bin/shard-init is missing.", "The background service is not running.") +} + +func TestUpgradeVerifiesBeforeItReplaces(t *testing.T) { + rs := newReleaseServer(t) + rs.add("v0.1.0", false, false, true, nil) + rs.add("v0.2.0", false, false, true, map[string]string{ + "shard-linux-amd64": "#!/bin/sh\necho client v0.2.0\n", + "shard-init-linux-amd64": "new init", + }) + f := newFakeHost(t) + m := linuxInstall("v0.1.0") + f.installed(t, m) + ui := &watchUI{fakeUI: confirming(true)} + ui.onPrint = func(lines []string) { + if slices.Contains(lines, "Latest release: v0.2.0") && rs.downloads.Load() != 0 { + t.Errorf("the version printed after %d downloads", rs.downloads.Load()) + } + } + + if err := (&Setup{Host: f.host(rs), UI: ui}).upgrade(t.Context(), m, ServiceActive); err != nil { + t.Fatalf("upgrade: %v", err) + } + for path, want := range map[string]string{ + "/usr/local/bin/shard": "#!/bin/sh\necho client v0.2.0\n", + "/usr/local/bin/shard-init": "new init", + "/home/u/.local/bin/shard": "#!/bin/sh\necho client v0.2.0\n", + } { + if got, _ := f.read(t, path); got != want { + t.Fatalf("%s = %q, want %q", path, got, want) + } + } + got, _, err := LoadManifest(f.host(nil)) + if err != nil || got.Version != "v0.2.0" || got.Provider != "gvisor" { + t.Fatalf("manifest = %+v, %v", got, err) + } + if !f.called("systemctl restart shard") || f.called("systemctl enable") { + t.Fatalf("calls = %v, want one restart and no unit change", f.calls) + } + said(t, ui.fakeUI, "Your sandboxes keep running while the daemon restarts.") +} + +func TestUpgradeKeepsTheOldBinaries(t *testing.T) { + cases := map[string]struct { + cli string + confirm bool + want error + wantMsg string + }{ + "wrong version": {cli: "#!/bin/sh\necho client v0.1.9\n", wantMsg: `reports "client v0.1.9"`}, + "declined": {cli: "#!/bin/sh\necho client v0.2.0\n", want: ErrDeclined}, + } + for name, tc := range cases { + t.Run(name, func(t *testing.T) { + rs := newReleaseServer(t) + rs.add("v0.2.0", false, false, true, map[string]string{"shard-linux-amd64": tc.cli, "shard-init-linux-amd64": "new init"}) + f := newFakeHost(t) + m := linuxInstall("v0.1.0") + f.installed(t, m) + + err := (&Setup{Host: f.host(rs), UI: confirming(tc.confirm)}).upgrade(t.Context(), m, ServiceActive) + if tc.want != nil && !errors.Is(err, tc.want) { + t.Fatalf("upgrade = %v, want %v", err, tc.want) + } + if tc.wantMsg != "" && (err == nil || !strings.Contains(err.Error(), tc.wantMsg)) { + t.Fatalf("upgrade = %v, want %q", err, tc.wantMsg) + } + if got, _ := f.read(t, "/usr/local/bin/shard"); got != "old shard" { + t.Fatalf("the old binary became %q", got) + } + if got, _, err := LoadManifest(f.host(nil)); err != nil || got.Version != "v0.1.0" { + t.Fatalf("manifest = %+v, %v", got, err) + } + if f.called("systemctl restart") { + t.Fatalf("calls = %v", f.calls) + } + }) + } +} + +func TestUpgradeStopsAtTheLatestRelease(t *testing.T) { + rs := newReleaseServer(t) + rs.add("v0.2.0", false, false, true, nil) + f := newFakeHost(t) + m := linuxInstall("v0.2.0") + f.installed(t, m) + ui := &fakeUI{} + + if err := (&Setup{Host: f.host(rs), UI: ui}).upgrade(t.Context(), m, ServiceActive); err != nil { + t.Fatalf("upgrade: %v", err) + } + said(t, ui, "Shard v0.2.0 is up to date.") + if len(ui.asked) != 0 || rs.downloads.Load() != 0 { + t.Fatalf("upgrade asked %v and downloaded %d files", ui.asked, rs.downloads.Load()) + } +} + +func TestUpgradeLeavesAnInactiveServiceStopped(t *testing.T) { + rs := newReleaseServer(t) + rs.add("v0.2.0", false, false, true, map[string]string{"shard-linux-amd64": "#!/bin/sh\necho client v0.2.0\n", "shard-init-linux-amd64": "new init"}) + f := newFakeHost(t) + m := linuxInstall("v0.1.0") + f.installed(t, m) + ui := confirming(true) + + if err := (&Setup{Host: f.host(rs), UI: ui}).upgrade(t.Context(), m, ServiceInactive); err != nil { + t.Fatalf("upgrade: %v", err) + } + said(t, ui, "Setup does not start it.") + if f.called("systemctl restart") || f.called("systemctl start") { + t.Fatalf("calls = %v", f.calls) + } +} + +func TestUninstallRefusesWhileASandboxRemains(t *testing.T) { + f := newFakeHost(t) + m := linuxInstall("v0.1.0") + f.installed(t, m) + f.write(t, "/var/lib/shard/sandboxes/sb_1/sandbox.json", "{}") + f.write(t, "/var/lib/shard/sandboxes/sb_2/sandbox.json", "{}") + f.write(t, "/var/lib/shard/images/sb_3/sandbox.json", "{}") + ui := &fakeUI{} + + err := (&Setup{Host: f.host(nil), UI: ui}).uninstall(t.Context(), m) + if err == nil || !strings.Contains(err.Error(), "2 sandboxes left") { + t.Fatalf("uninstall = %v, want 2 sandboxes named", err) + } + said(t, ui, "Shard has 2 sandboxes on this machine.", "shard list --all", "shard remove --force <name>") + if _, ok := f.read(t, "/usr/local/bin/shard"); !ok || len(ui.asked) != 0 { + t.Fatalf("uninstall removed a file or asked %v while a sandbox remains", ui.asked) + } +} + +func TestUninstallRemovesOnlyWhatSetupOwns(t *testing.T) { + f := newFakeHost(t) + m := linuxInstall("v0.1.0") + f.installed(t, m) + f.write(t, "/var/lib/shard/images/index.json", "{}") + ui := &watchUI{fakeUI: confirming(true)} + + if err := (&Setup{Host: f.host(nil), UI: ui}).uninstall(t.Context(), m); err != nil { + t.Fatalf("uninstall: %v", err) + } + if !slices.Equal(ui.defaults, []bool{false}) { + t.Fatalf("Confirm defaults = %v, want one No", ui.defaults) + } + for _, gone := range []string{"/usr/local/bin/shard", "/usr/local/bin/shard-init", "/etc/systemd/system/shard.service", "/etc/shard/serve.env", ManifestPath, filepath.Dir(ManifestPath)} { + if _, err := os.Lstat(filepath.Join(f.root, gone)); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("%s is still there: %v", gone, err) + } + } + for _, kept := range []string{"/usr/sbin/nft", "/var/lib/shard/images/index.json"} { + if _, ok := f.read(t, kept); !ok { + t.Fatalf("uninstall removed %s", kept) + } + } + if !f.called("systemctl disable --now shard") || !f.called("systemctl daemon-reload") { + t.Fatalf("calls = %v", f.calls) + } + said(t, ui.fakeUI, "Your saved data remains in /var/lib/shard.", "sudo apt-get remove nftables", "rm "+f.host(nil).Executable) +} + +func TestUninstallDeclinedChangesNothing(t *testing.T) { + f := newFakeHost(t) + m := linuxInstall("v0.1.0") + f.installed(t, m) + ui := confirming(false) + + err := (&Setup{Host: f.host(nil), UI: ui}).uninstall(t.Context(), m) + if !errors.Is(err, ErrDeclined) { + t.Fatalf("uninstall = %v, want ErrDeclined", err) + } + want := []string{"Uninstall Shard?", "", + "This will stop and remove the background service", + "and remove files installed by Shard setup.", "", + "Your saved data will remain.", + "Shared tools will remain.", ""} + if !slices.Equal(ui.printed, want) { + t.Fatalf("output = %q, want %q", ui.printed, want) + } + if _, ok, err := LoadManifest(f.host(nil)); err != nil || !ok || len(f.calls) != 0 { + t.Fatalf("a declined uninstall ran %v (manifest %v, %v)", f.calls, ok, err) + } +} + +func TestUninstallRunsAgainAfterAHalfDoneOne(t *testing.T) { + f := newFakeHost(t) + m := linuxInstall("v0.1.0") + f.installed(t, m) + if err := os.Remove(filepath.Join(f.root, "/etc/systemd/system/shard.service")); err != nil { + t.Fatalf("remove: %v", err) + } + + if err := (&Setup{Host: f.host(nil), UI: confirming(true)}).uninstall(t.Context(), m); err != nil { + t.Fatalf("uninstall: %v", err) + } + if f.called("systemctl disable") { + t.Fatalf("uninstall stopped a unit that is gone: %v", f.calls) + } + if _, ok, err := LoadManifest(f.host(nil)); err != nil || ok { + t.Fatalf("the manifest outlived the uninstall: %v, %v", ok, err) + } +} + +// A repair is local setup too, so it offers to drop a saved remote and drops it once the repair succeeds; Exit asks nothing. +func TestTheExistingMenuOffersToDropASavedRemote(t *testing.T) { + saved := client.Config{Remote: "https://shard.example.com", APIKey: testKey} + + for _, tc := range []struct { + choice string + asks bool + want client.Config + }{ + {choice: "repair", asks: true, want: client.Config{}}, + {choice: "exit", want: saved}, + } { + t.Run(tc.choice, func(t *testing.T) { + f := newFakeHost(t) + m := linuxInstall("v0.1.0") + f.installed(t, m) + env, path := testHost(t, nil) + saveConnection(t, path, saved) + host := f.host(nil) + host.Env = env.Env + ui := &fakeUI{selects: map[Question]string{AskExisting: tc.choice}, confirms: map[Question]bool{AskSwitch: true}} + + if err := (&Setup{Host: host, UI: ui}).existing(t.Context(), Installation{Manifest: &m, Service: ServiceActive}); err != nil { + t.Fatalf("existing: %v", err) + } + if got := slices.Contains(ui.asked, AskSwitch); got != tc.asks { + t.Errorf("asked %v, want the switch asked: %v", ui.asked, tc.asks) + } + if got := savedConnection(t, path); got != tc.want { + t.Errorf("the saved connection is %v, want %v", got, tc.want) + } + }) + } +} diff --git a/services/setup/local.go b/services/setup/local.go new file mode 100644 index 000000000..03635b6d4 --- /dev/null +++ b/services/setup/local.go @@ -0,0 +1,199 @@ +package setup + +import ( + "context" + "errors" + "fmt" + "io/fs" + "os" + "strings" + + "github.com/presmihaylov/shard/pkg/term" +) + +// Local is what a local setup installs: one provider, and whether a service starts it at boot. +type Local struct { + Provider string + StartAtBoot bool +} + +// local is §5 to §10: the provider, automatic startup, preflight, review, and apply. +func (s *Setup) local(ctx context.Context) error { + provider, err := s.askProvider(ctx, nil) + if err != nil { + return err + } + startAtBoot, err := s.askStartAtBoot(ctx) + if err != nil { + return err + } + l, err := s.checked(ctx, Local{Provider: provider, StartAtBoot: startAtBoot}) + if err != nil { + return err + } + if err := s.review(ctx, l); err != nil { + return err + } + if err := s.admin(ctx); err != nil { + return err + } + steps, err := s.localSteps(ctx, l) + if err != nil { + return err + } + if err := s.apply(ctx, "Setting up Shard", steps); err != nil { + return err + } + + return s.UI.Print(localDone(s.Host, l)...) +} + +const exitOption = "exit" + +// askProvider is §5; a provider that failed preflight shows why, and Exit leaves with nothing changed. +func (s *Setup) askProvider(ctx context.Context, failed map[string][]string) (string, error) { + choices := Providers(ctx, s.Host) + options := make([]term.Option, 0, len(choices)+1) + available := false + for _, c := range choices { + o := term.Option{Name: c.Name, Label: c.Title, Lines: c.Lines, Unavailable: c.Unavailable} + if lines, ok := failed[c.Name]; ok { + o.Lines, o.Unavailable = nil, lines + } + o.Default = c.Recommended && len(o.Unavailable) == 0 + available = available || len(o.Unavailable) == 0 + options = append(options, o) + } + // The menu always keeps one row to choose, since a list with none refuses to draw. + if len(failed) > 0 || !available { + options = append(options, term.Option{Name: exitOption, Label: "Exit"}) + } + i, err := s.UI.Select(ctx, AskProvider, "Choose how Shard isolates your sandboxes:", options) + if err != nil { + return "", err + } + if options[i].Name == exitOption { + return "", ErrDeclined + } + + return options[i].Name, nil +} + +// askStartAtBoot is §7; the names are what --start-at-boot takes. +func (s *Setup) askStartAtBoot(ctx context.Context) (bool, error) { + options := []term.Option{ + {Name: "true", Label: "Yes (recommended)", Default: true, Lines: []string{ + "Set up a background service using systemd on Linux or launchd on macOS.", + "The service starts Shard at boot and restarts it after a crash.", + }}, + {Name: "false", Label: "No", Lines: []string{ + "Manage the Shard daemon yourself.", + "Local sandbox commands fail when the daemon is not running.", + "Run `shard daemon` in a terminal, or configure your own background", + "service and automatic startup.", + }}, + } + i, err := s.UI.Select(ctx, AskStartAtBoot, "Start Shard automatically when this machine starts?", options) + if err != nil { + return false, err + } + + return options[i].Name == "true", nil +} + +// checked runs preflight until it passes; a failure another provider may not have asks for the provider again. +func (s *Setup) checked(ctx context.Context, l Local) (Local, error) { + failed := map[string][]string{} + for { + f, err := s.preflight(ctx, l) + if err != nil || f == nil { + return l, err + } + if err := s.UI.Print("", "No installation changes were made."); err != nil { + return l, err + } + if !f.provider { + return l, &StoppedError{Step: f.check, Err: &Problem{Lines: f.lines}} + } + if err := s.UI.Print("Choose another provider or exit.", ""); err != nil { + return l, err + } + failed[l.Provider] = f.lines + if l.Provider, err = s.askProvider(ctx, failed); err != nil { + return l, err + } + } +} + +// review is §9: the changes setup will make, and the confirmation before any of them. +func (s *Setup) review(ctx context.Context, l Local) error { + title := providerTitle(l.Provider) + startup := "No" + if l.StartAtBoot { + startup = "Yes" + } + lines := []string{"Ready to set up Shard", "", "Provider: " + title, "Automatic startup: " + startup, "", "Setup will:"} + if names := missingTools(s.Host, l.Provider).names(); len(names) > 0 { + lines = append(lines, fmt.Sprintf(" Install the tools required by %s: %s.", title, strings.Join(names, ", "))) + } + if binaries := shardBinaries(s.Host); len(binaries) > 0 { + lines = append(lines, " Install "+strings.Join(binaries, " and ")+" in "+binDir+".") + } + _, err := os.Lstat(rooted(s.Host, DataDir)) + switch { + case errors.Is(err, fs.ErrNotExist): + lines = append(lines, " Create Shard's data directory.") + case err != nil: + return fmt.Errorf("check %s: %w", DataDir, err) + } + lines = append(lines, startupLine(s.Host, l), "", "Administrator access is required.", "") + if err := s.UI.Print(lines...); err != nil { + return err + } + + return s.confirm(ctx, true) +} + +// shardBinaries are what installShard puts in binDir: shard unless it already runs from there, and shard-init on Linux. +func shardBinaries(h Host) []string { + var names []string + if h.Executable != shardBinary { + names = append(names, "shard") + } + if h.OS == "linux" { + names = append(names, "shard-init") + } + + return names +} + +func startupLine(h Host, l Local) string { + switch { + case !l.StartAtBoot: + return " Leave daemon startup under your control." + case h.OS == "darwin": + return " Configure and start a launchd service." + } + + return " Configure and start a systemd service." +} + +// localDone closes a local setup with what to run next; on Linux the socket belongs to root, so local commands need sudo. +func localDone(h Host, l Local) []string { + sudo := "" + lines := []string{""} + if h.OS == "linux" { + sudo = "sudo " + } + if l.StartAtBoot { + lines = append(lines, "Shard is set up, and the daemon is running.") + } + if !l.StartAtBoot { + lines = append(lines, "Shard is set up.", "Start the daemon with:", " "+sudo+"shard daemon --provider "+l.Provider, "") + } + if sudo != "" { + lines = append(lines, "Local commands run with sudo, for example `sudo shard ls`.") + } + + return lines +} diff --git a/services/setup/local_test.go b/services/setup/local_test.go new file mode 100644 index 000000000..f534d4a2a --- /dev/null +++ b/services/setup/local_test.go @@ -0,0 +1,489 @@ +package setup + +import ( + "archive/tar" + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "fmt" + "maps" + "os" + "os/exec" + "path/filepath" + "slices" + "strings" + "testing" + + "github.com/klauspost/compress/zstd" + + "github.com/presmihaylov/shard/models" + "github.com/presmihaylov/shard/pkg/term" + "github.com/presmihaylov/shard/services/daemon" + "github.com/presmihaylov/shard/services/sandboxstate" +) + +// localHost is a Linux host under a temp root with gVisor's tools, systemd and a release in place, whose commands a test records. +type localHost struct { + t *testing.T + root string + rs *releaseServer + env map[string]string + calls []string + // fail makes a command whose line starts with a key fail, printing its value. + fail map[string]string +} + +func newLocalHost(t *testing.T) *localHost { + t.Helper() + swap(t, &rootUID, os.Getuid()) + l := &localHost{t: t, root: t.TempDir(), rs: newReleaseServer(t), env: map[string]string{}, fail: map[string]string{}} + l.rs.add("v0.1.0", false, false, true, map[string]string{"shard-init-linux-amd64": "shard-init v0.1.0"}) + for _, tool := range []string{"/usr/sbin/ip", "/usr/sbin/nft", "/usr/sbin/mkfs.ext4", "/usr/local/bin/runsc", "/usr/bin/systemctl", "/usr/bin/systemd-analyze", "/usr/bin/apt-get", "/usr/bin/sudo"} { + l.write(tool, "#!/bin/sh\n") + } + l.write("/home/u/.local/bin/shard", "shard v0.1.0") + l.mkdir("/run/systemd/system") + + return l +} + +func (l *localHost) host() Host { + return Host{ + Root: l.root, OS: "linux", Arch: "amd64", Version: "v0.1.0", + Executable: filepath.Join(l.root, "/home/u/.local/bin/shard"), + Releases: l.rs.URL + "/releases", HTTP: l.rs.Client(), + Env: func(k string) string { return l.env[k] }, + Run: l.run, + } +} + +func (l *localHost) mac() Host { + h := l.host() + h.OS, h.Arch, h.Euid = "darwin", "arm64", os.Getuid() + l.env["USER"] = "u" + for _, tool := range []string{"/usr/bin/codesign", "/bin/launchctl", "/usr/bin/plutil"} { + l.write(tool, "#!/bin/sh\n") + } + + return h +} + +// run does the file work for real, minus the chown a test cannot make, and answers every other command with success. +func (l *localHost) run(ctx context.Context, name string, args ...string) ([]byte, error) { + line := strings.Join(append([]string{name}, args...), " ") + l.calls = append(l.calls, line) + for prefix, out := range l.fail { + if strings.HasPrefix(line, prefix) { + return []byte(out), errors.New("exit status 1") + } + } + + return l.do(ctx, name, args...) +} + +func (l *localHost) do(ctx context.Context, name string, args ...string) ([]byte, error) { + switch name { + case "sudo": + if len(args) > 2 && args[0] == "-n" && args[1] == "--" { + return l.do(ctx, args[2], args[3:]...) + } + return nil, nil + case "install": + return exec.CommandContext(ctx, name, withoutOwner(args)...).CombinedOutput() + case "mkdir", "mv", "rm", "rmdir": + return exec.CommandContext(ctx, name, args...).CombinedOutput() + case "sw_vers": + return []byte("14.5\n"), nil + } + if strings.HasPrefix(name, l.root) { + return exec.CommandContext(ctx, name, args...).CombinedOutput() + } + + return nil, nil +} + +func withoutOwner(args []string) []string { + var kept []string + for i := 0; i < len(args); i++ { + if args[i] == "-o" || args[i] == "-g" { + i++ + continue + } + kept = append(kept, args[i]) + } + + return kept +} + +func (l *localHost) write(path, body string) { + l.t.Helper() + full := filepath.Join(l.root, path) + if err := os.MkdirAll(filepath.Dir(full), 0o755); err != nil { + l.t.Fatalf("mkdir %s: %v", path, err) + } + if err := os.WriteFile(full, []byte(body), 0o755); err != nil { + l.t.Fatalf("write %s: %v", path, err) + } +} + +func (l *localHost) mkdir(path string) { + l.t.Helper() + if err := os.MkdirAll(filepath.Join(l.root, path), 0o755); err != nil { + l.t.Fatalf("mkdir %s: %v", path, err) + } +} + +func (l *localHost) remove(path string) { + l.t.Helper() + if err := os.Remove(filepath.Join(l.root, path)); err != nil { + l.t.Fatalf("remove %s: %v", path, err) + } +} + +func (l *localHost) exists(path string) bool { + _, err := os.Lstat(filepath.Join(l.root, path)) + if err != nil && !errors.Is(err, os.ErrNotExist) { + l.t.Fatalf("stat %s: %v", path, err) + } + + return err == nil +} + +// changed is whether setup ran any command that changes the host. +func (l *localHost) changed() bool { + return slices.ContainsFunc(l.calls, func(c string) bool { + return !strings.HasPrefix(c, "sw_vers") && !strings.HasSuffix(c, "--version") + }) +} + +// sandbox records one sandbox of provider in the data dir. +func (l *localHost) sandbox(provider string) { + l.t.Helper() + repo, err := sandboxstate.New(filepath.Join(l.root, DataDir)) + if err != nil { + l.t.Fatalf("open the sandbox records: %v", err) + } + if _, err := repo.Create(models.Sandbox{Provider: provider, State: models.StateRunning}); err != nil { + l.t.Fatalf("record a sandbox: %v", err) + } +} + +// pinGVisor serves a gVisor release from the fake server and drops runsc, so setup downloads and installs it. +func (l *localHost) pinGVisor() { + l.t.Helper() + files := map[string]string{"runsc": "/usr/local/bin/runsc", "gvisor-bin/gvisor_sentry": "/usr/local/bin/gvisor-bin/gvisor_sentry"} + archive := tarZstd(l.t, files) + sum := sha256.Sum256(archive) + l.rs.files["pins/gvisor.tar.zstd"] = string(archive) + swap(l.t, gvisorRelease, download{Title: "runsc", URL: l.rs.URL + "/download/pins/gvisor.tar.zstd", SHA256: hex.EncodeToString(sum[:]), Files: files}) + l.remove("/usr/local/bin/runsc") +} + +func tarZstd(t *testing.T, files map[string]string) []byte { + t.Helper() + var buf bytes.Buffer + zw, err := zstd.NewWriter(&buf) + if err != nil { + t.Fatalf("zstd: %v", err) + } + tw := tar.NewWriter(zw) + for _, name := range slices.Sorted(maps.Keys(files)) { + body := "binary " + name + if err := tw.WriteHeader(&tar.Header{Name: name, Mode: 0o755, Size: int64(len(body)), Typeflag: tar.TypeReg}); err != nil { + t.Fatalf("tar header: %v", err) + } + if _, err := tw.Write([]byte(body)); err != nil { + t.Fatalf("tar body: %v", err) + } + } + if err := errors.Join(tw.Close(), zw.Close()); err != nil { + t.Fatalf("close the archive: %v", err) + } + + return buf.Bytes() +} + +func swap[T any](t *testing.T, p *T, v T) { + t.Helper() + old := *p + *p = v + t.Cleanup(func() { *p = old }) +} + +// localUI answers the provider question from a queue, so a test can choose again after a failed preflight. +type localUI struct { + *fakeUI + providers []string + shown [][]term.Option +} + +func (u *localUI) Select(ctx context.Context, q Question, title string, options []term.Option) (int, error) { + if q == AskProvider { + u.shown = append(u.shown, options) + u.selects[q], u.providers = u.providers[0], u.providers[1:] + } + + return u.fakeUI.Select(ctx, q, title, options) +} + +func newLocalUI(startAtBoot string, confirm bool, providers ...string) *localUI { + return &localUI{ + fakeUI: &fakeUI{selects: map[Question]string{AskStartAtBoot: startAtBoot}, confirms: map[Question]bool{AskConfirm: confirm}}, + providers: providers, + } +} + +func marks(list *fakeChecklist) string { return strings.Join(list.marks, ", ") } + +func finished(t *testing.T, list *fakeChecklist) { + t.Helper() + for i := range list.steps { + if !slices.Contains(list.marks, fmt.Sprintf("done %d", i)) { + t.Errorf("%s: step %d %q is not done: %s", list.title, i, list.steps[i], marks(list)) + } + } +} + +func TestLocalSetsUpGVisorWithAService(t *testing.T) { + l := newLocalHost(t) + l.pinGVisor() + ui := newLocalUI("true", true, GVisor) + + if err := (&Setup{Host: l.host(), UI: ui}).local(t.Context()); err != nil { + t.Fatalf("local = %v; printed %q", err, ui.printed) + } + + if want := []Question{AskProvider, AskStartAtBoot, AskConfirm}; !slices.Equal(ui.asked, want) { + t.Fatalf("asked %v, want %v", ui.asked, want) + } + said(t, ui.fakeUI, + "Provider: gVisor", + "Automatic startup: Yes", + " Install the tools required by gVisor: runsc.", + " Install shard and shard-init in /usr/local/bin.", + " Create Shard's data directory.", + " Configure and start a systemd service.", + "Administrator access is required.", + "Shard is set up, and the daemon is running.", + "Local commands run with sudo, for example `sudo shard ls`.", + ) + if len(ui.lists) != 2 || ui.lists[0].title != "Checking this machine" || ui.lists[1].title != "Setting up Shard" { + t.Fatalf("checklists %v", ui.lists) + } + finished(t, ui.lists[0]) + finished(t, ui.lists[1]) + for _, path := range []string{"/usr/local/bin/runsc", "/usr/local/bin/gvisor-bin/gvisor_sentry", shardBinary, initBinary, DataDir, systemdUnit, ManifestPath} { + if !l.exists(path) { + t.Errorf("%s is not installed", path) + } + } + unit, err := os.ReadFile(filepath.Join(l.root, systemdUnit)) + if err != nil || !strings.Contains(string(unit), "ExecStart=/usr/local/bin/shard daemon --provider gvisor") { + t.Fatalf("the unit names no provider: %v\n%s", err, unit) + } + if !slices.Contains(l.calls, "systemctl start shard.service") { + t.Fatalf("the daemon never started: %q", l.calls) + } +} + +func TestLocalManualStartupPrintsTheDaemonCommand(t *testing.T) { + l := newLocalHost(t) + ui := newLocalUI("false", true, GVisor) + + if err := (&Setup{Host: l.host(), UI: ui}).local(t.Context()); err != nil { + t.Fatalf("local = %v; printed %q", err, ui.printed) + } + + said(t, ui.fakeUI, "Automatic startup: No", " Leave daemon startup under your control.", "Start the daemon with:", " sudo shard daemon --provider gvisor") + if slices.ContainsFunc(ui.printed, func(p string) bool { return strings.Contains(p, "Install the tools") }) { + t.Fatalf("the review offers tools the host has: %q", ui.printed) + } + if len(ui.lists[0].steps) != 7 || slices.Contains(ui.lists[0].steps, "Background service support") { + t.Fatalf("manual startup checks %q", ui.lists[0].steps) + } + if slices.ContainsFunc(l.calls, func(c string) bool { return strings.HasPrefix(c, "systemctl") }) || l.exists(systemdUnit) { + t.Fatalf("manual startup touched systemd: %q", l.calls) + } +} + +func TestLocalOnAMacUsesLaunchdAndNoSudo(t *testing.T) { + l := newLocalHost(t) + h := l.mac() + swap(t, &kernelURL, func(string) (string, error) { return l.rs.URL + "/download/v0.1.0/shard-init-linux-amd64", nil }) + ui := newLocalUI("true", true, VZ) + + if err := (&Setup{Host: h, UI: ui}).local(t.Context()); err != nil { + t.Fatalf("local = %v; printed %q", err, ui.printed) + } + + said(t, ui.fakeUI, "Provider: macOS Virtualization", " Install shard in /usr/local/bin.", " Configure and start a launchd service.", "Shard is set up, and the daemon is running.") + if slices.ContainsFunc(ui.printed, func(p string) bool { return strings.Contains(p, "sudo") }) { + t.Fatalf("a Mac is told to use sudo: %q", ui.printed) + } + plist, err := os.ReadFile(filepath.Join(l.root, launchdPlist)) + if err != nil || !strings.Contains(string(plist), "<string>u</string>") { + t.Fatalf("the plist names no user: %v\n%s", err, plist) + } + if !l.exists(shardBinary) || !l.exists(DataDir) || l.exists(initBinary) { + t.Fatal("a Mac needs shard and the data dir, and never shard-init, which its daemon embeds") + } +} + +func TestADeclinedReviewChangesNothing(t *testing.T) { + l := newLocalHost(t) + l.pinGVisor() + ui := newLocalUI("true", false, GVisor) + + err := (&Setup{Host: l.host(), UI: ui}).local(t.Context()) + if !errors.Is(err, ErrDeclined) { + t.Fatalf("local = %v, want ErrDeclined", err) + } + if l.changed() || l.exists(shardBinary) || l.exists(DataDir) { + t.Fatalf("a declined setup changed the host: %q", l.calls) + } +} + +func TestAProviderFailureOffersTheOthers(t *testing.T) { + l := newLocalHost(t) + l.sandbox(GVisor) + ui := newLocalUI("true", false, Runc, GVisor) + + err := (&Setup{Host: l.host(), UI: ui}).local(t.Context()) + if !errors.Is(err, ErrDeclined) { + t.Fatalf("local = %v, want the review after the second choice", err) + } + + said(t, ui.fakeUI, "No installation changes were made.", "Choose another provider or exit.", "Provider: gVisor") + if len(ui.shown) != 2 { + t.Fatalf("asked for the provider %d times, want 2", len(ui.shown)) + } + again := ui.shown[1] + runc := again[providerIndex(Runc)] + if !slices.Equal(runc.Unavailable, []string{"The sandboxes in /var/lib/shard use gVisor.", "Setup never changes the provider of existing sandboxes."}) { + t.Fatalf("the failed provider shows %+v", runc) + } + if last := again[len(again)-1]; last.Name != exitOption { + t.Fatalf("the second question has no Exit: %+v", last) + } + if want := []Question{AskProvider, AskStartAtBoot, AskProvider, AskConfirm}; !slices.Equal(ui.asked, want) { + t.Fatalf("asked %v, want %v", ui.asked, want) + } + if l.changed() { + t.Fatalf("a failed preflight changed the host: %q", l.calls) + } +} + +func TestExitAfterAProviderFailureChangesNothing(t *testing.T) { + l := newLocalHost(t) + l.sandbox(Runc) + ui := newLocalUI("true", true, GVisor, exitOption) + + err := (&Setup{Host: l.host(), UI: ui}).local(t.Context()) + if !errors.Is(err, ErrDeclined) { + t.Fatalf("local = %v, want ErrDeclined", err) + } + if slices.Contains(ui.asked, AskConfirm) || l.changed() { + t.Fatalf("exit went on: asked %v, ran %q", ui.asked, l.calls) + } +} + +func TestAHostFailureStopsWithoutAnotherChoice(t *testing.T) { + l := newLocalHost(t) + l.remove("/run/systemd/system") + ui := newLocalUI("true", true, GVisor) + + err := (&Setup{Host: l.host(), UI: ui}).local(t.Context()) + var stopped *StoppedError + if !errors.As(err, &stopped) || stopped.Step != "Background service support" { + t.Fatalf("local = %v, want a stop at the service check", err) + } + said(t, ui.fakeUI, "No installation changes were made.") + if slices.Contains(ui.printed, "Choose another provider or exit.") || len(ui.shown) != 1 { + t.Fatalf("a host failure offered another provider: %q", ui.printed) + } + want := "fail 7: Setup configures a systemd service, and systemd does not manage this machine. / Run shard setup again and choose No for automatic startup." + if got := ui.lists[0].marks[len(ui.lists[0].marks)-1]; got != want { + t.Fatalf("last mark %q, want %q", got, want) + } +} + +func TestEveryProviderUnavailableStillOffersExit(t *testing.T) { + l := newLocalHost(t) + h := l.host() + h.OS = "darwin" + ui := newLocalUI("true", true, exitOption) + + err := (&Setup{Host: h, UI: ui}).local(t.Context()) + if !errors.Is(err, ErrDeclined) { + t.Fatalf("local = %v, want ErrDeclined", err) + } + options := ui.shown[0] + for _, o := range options[:len(options)-1] { + if len(o.Unavailable) == 0 { + t.Errorf("%s is available on an Intel Mac", o.Name) + } + } + if options[len(options)-1].Name != exitOption { + t.Fatalf("no Exit among %+v", options) + } +} + +func TestAFailedStepSaysWhatStays(t *testing.T) { + l := newLocalHost(t) + l.fail["systemctl start"] = "Job for shard.service failed." + ui := newLocalUI("true", true, GVisor) + + err := (&Setup{Host: l.host(), UI: ui}).local(t.Context()) + var stopped *StoppedError + if !errors.As(err, &stopped) || stopped.Step != "Start the daemon" { + t.Fatalf("local = %v, want a stop at the start", err) + } + said(t, ui.fakeUI, "Setup stopped. Earlier completed steps remain in place.") + if !l.exists(systemdUnit) { + t.Fatal("the stop undid the service it had installed") + } +} + +// The templates are copies of the packaged files, so a change to one must reach the other. +func TestServiceTemplatesMatchPackaging(t *testing.T) { + for file, template := range map[string]string{ + "../../packaging/systemd/shard.service": systemdUnitTemplate, + "../../packaging/launchd/shard.daemon.plist": launchdPlistTemplate, + "../../packaging/launchd/shard.newsyslog.conf": newsyslogTemplate, + } { + packaged, err := os.ReadFile(file) + if err != nil { + t.Fatalf("read %s: %v", file, err) + } + if string(packaged) != template { + t.Errorf("%s differs from its template in apply_local.go", file) + } + } +} + +// Setup writes the provider into the daemon's command line, so its names must be the ones the daemon takes. +func TestProvidersAreTheDaemons(t *testing.T) { + var names []string + for _, p := range providerTexts { + names = append(names, p.name) + } + if !slices.Equal(slices.Sorted(slices.Values(names)), slices.Sorted(slices.Values(daemon.Providers))) { + t.Fatalf("setup offers %q, and the daemon takes %q", names, daemon.Providers) + } +} + +// A failed verify says what the daemon reported, never the sudo command line that asked it. +func TestNotReadyIsTheDaemonsReason(t *testing.T) { + err := errors.New("sudo -n -- /usr/local/bin/shard --remote daemon status: exit status 1") + for out, want := range map[string]string{ + "shard: tasks in backoff: egress-log-tailer\n": "tasks in backoff: egress-log-tailer", + "": err.Error(), + " \n\n ": err.Error(), + } { + if got := notReady([]byte(out), err); got != want { + t.Errorf("notReady(%q) = %q, want %q", out, got, want) + } + } +} diff --git a/services/setup/manifest.go b/services/setup/manifest.go new file mode 100644 index 000000000..0176d5880 --- /dev/null +++ b/services/setup/manifest.go @@ -0,0 +1,153 @@ +package setup + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "os" + "path/filepath" + "slices" +) + +// ManifestPath sits outside the data dir and is root-owned, since uninstall removes with root what it names. +const ManifestPath = "/var/lib/shard-setup/manifest.json" + +// Kind says what uninstall does with an Owned file. +type Kind string + +const ( + // KindBinary, KindService and KindConfig are shard's own; uninstall removes them. + KindBinary Kind = "binary" + KindService Kind = "service" + KindConfig Kind = "config" + // KindTool is a shared dependency setup installed; uninstall lists it and keeps it. + KindTool Kind = "tool" + // KindData is saved data; uninstall keeps it. + KindData Kind = "data" +) + +// Owned is one file setup created. Path is on the host, without Host.Root. +type Owned struct { + Path string `json:"path"` + Kind Kind `json:"kind"` + // Package is the system package a tool came from, so uninstall can name the command that removes it. + Package string `json:"package,omitempty"` +} + +// Manifest is what setup installed, in the order it did it. +type Manifest struct { + Version string `json:"version"` + Provider string `json:"provider"` + StartAtBoot bool `json:"start_at_boot"` + Files []Owned `json:"files"` +} + +// RecordOwned adds files to the manifest after setup created them; a file that was already there is never recorded. +func RecordOwned(ctx context.Context, h Host, local Local, files ...Owned) error { + m, _, err := LoadManifest(h) + if err != nil { + return err + } + + m.Version, m.Provider, m.StartAtBoot = h.Version, local.Provider, local.StartAtBoot + for _, f := range files { + if err := f.valid(); err != nil { + return err + } + i := slices.IndexFunc(m.Files, func(o Owned) bool { return o.Path == f.Path }) + if i >= 0 { + m.Files[i] = f + continue + } + m.Files = append(m.Files, f) + } + + return saveManifest(ctx, h, m) +} + +// LoadManifest reads the manifest, and reports false when setup never installed on this host. +func LoadManifest(h Host) (Manifest, bool, error) { + path := filepath.Join(h.Root, ManifestPath) + data, err := os.ReadFile(path) + if errors.Is(err, os.ErrNotExist) { + return Manifest{}, false, nil + } + if err != nil { + return Manifest{}, false, fmt.Errorf("read the installation manifest: %w", err) + } + + var m Manifest + if err := json.Unmarshal(data, &m); err != nil { + return Manifest{}, false, fmt.Errorf("decode the installation manifest %s: %w", path, err) + } + for _, f := range m.Files { + if err := f.valid(); err != nil { + return Manifest{}, false, fmt.Errorf("installation manifest %s: %w", path, err) + } + } + + return m, true, nil +} + +// saveManifest writes as the user, then moves the file into the root-owned dir with one rename. +func saveManifest(ctx context.Context, h Host, m Manifest) error { + data, err := json.MarshalIndent(m, "", " ") + if err != nil { + return fmt.Errorf("encode the installation manifest: %w", err) + } + + tmp, err := os.CreateTemp("", "shard-manifest-*.json") + if err != nil { + return fmt.Errorf("write the installation manifest: %w", err) + } + _, err = tmp.Write(append(data, '\n')) + if err := errors.Join(err, tmp.Close()); err != nil { + return errors.Join(fmt.Errorf("write the installation manifest: %w", err), os.Remove(tmp.Name())) + } + + path := filepath.Join(h.Root, ManifestPath) + next := path + ".new" + steps := [][]string{ + {"mkdir", "-p", "-m", "0755", filepath.Dir(path)}, + {"install", "-m", "0644", tmp.Name(), next}, + {"mv", "-f", next, path}, + } + for _, step := range steps { + if _, err := privileged(ctx, h, step[0], step[1:]...); err != nil { + return errors.Join(fmt.Errorf("save the installation manifest: %w", err), os.Remove(tmp.Name())) + } + } + + if err := os.Remove(tmp.Name()); err != nil { + return fmt.Errorf("save the installation manifest: %w", err) + } + + return nil +} + +// removeManifest is the last step of an uninstall, so a failed one can run again from the same list. +func removeManifest(ctx context.Context, h Host) error { + path := filepath.Join(h.Root, ManifestPath) + if _, err := privileged(ctx, h, "rm", "-f", path); err != nil { + return fmt.Errorf("remove the installation manifest: %w", err) + } + if _, err := privileged(ctx, h, "rmdir", filepath.Dir(path)); err != nil { + return fmt.Errorf("remove %s: %w", filepath.Dir(path), err) + } + + return nil +} + +// valid keeps a manifest from naming a relative or unclean path, which a root rm would resolve somewhere else. +func (f Owned) valid() error { + if !filepath.IsAbs(f.Path) || filepath.Clean(f.Path) != f.Path || f.Path == "/" { + return fmt.Errorf("owned file %q is not a clean absolute path", f.Path) + } + switch f.Kind { + case KindBinary, KindService, KindConfig, KindTool, KindData: + return nil + } + + return fmt.Errorf("owned file %s has unknown kind %q", f.Path, f.Kind) +} diff --git a/services/setup/manifest_test.go b/services/setup/manifest_test.go new file mode 100644 index 000000000..1efcb7bac --- /dev/null +++ b/services/setup/manifest_test.go @@ -0,0 +1,80 @@ +package setup + +import ( + "errors" + "os" + "path/filepath" + "slices" + "strings" + "testing" +) + +func TestRecordOwnedMergesByPath(t *testing.T) { + f := newFakeHost(t) + h := f.host(nil) + local := Local{Provider: "gvisor", StartAtBoot: true} + + if err := RecordOwned(t.Context(), h, local, Owned{Path: "/usr/sbin/nft", Kind: KindTool}, Owned{Path: "/usr/local/bin/shard", Kind: KindBinary}); err != nil { + t.Fatalf("RecordOwned: %v", err) + } + h.Version = "v0.2.0" + if err := RecordOwned(t.Context(), h, Local{Provider: "runc"}, Owned{Path: "/usr/sbin/nft", Kind: KindTool, Package: "nftables"}); err != nil { + t.Fatalf("RecordOwned: %v", err) + } + + m, ok, err := LoadManifest(h) + if err != nil || !ok { + t.Fatalf("LoadManifest = %v, %v", ok, err) + } + want := []Owned{{Path: "/usr/sbin/nft", Kind: KindTool, Package: "nftables"}, {Path: "/usr/local/bin/shard", Kind: KindBinary}} + if !slices.Equal(m.Files, want) || m.Version != "v0.2.0" || m.Provider != "runc" || m.StartAtBoot { + t.Fatalf("manifest = %+v, want %v at v0.2.0 on runc without start at boot", m, want) + } +} + +func TestRecordOwnedRefusesAnUncleanPath(t *testing.T) { + f := newFakeHost(t) + + err := RecordOwned(t.Context(), f.host(nil), Local{}, Owned{Path: "/usr/local/../bin/shard", Kind: KindBinary}) + if err == nil || !strings.Contains(err.Error(), "not a clean absolute path") { + t.Fatalf("RecordOwned = %v, want the path refused", err) + } + if _, ok, err := LoadManifest(f.host(nil)); err != nil || ok { + t.Fatalf("a refused record wrote a manifest: %v, %v", ok, err) + } +} + +func TestLoadManifestRefusesWhatARootRemoveCouldMisread(t *testing.T) { + cases := map[string]struct { + body string + want string + }{ + "relative path": {body: `{"files":[{"path":"usr/local/bin/shard","kind":"binary"}]}`, want: "not a clean absolute path"}, + "the root": {body: `{"files":[{"path":"/","kind":"binary"}]}`, want: "not a clean absolute path"}, + "unknown kind": {body: `{"files":[{"path":"/usr/local/bin/shard","kind":"cache"}]}`, want: `unknown kind "cache"`}, + "not json": {body: `{`, want: "decode the installation manifest"}, + } + for name, tc := range cases { + t.Run(name, func(t *testing.T) { + f := newFakeHost(t) + f.write(t, ManifestPath, tc.body) + + _, _, err := LoadManifest(f.host(nil)) + if err == nil || !strings.Contains(err.Error(), tc.want) { + t.Fatalf("LoadManifest = %v, want %q", err, tc.want) + } + }) + } +} + +func TestRemoveManifestTakesItsDirectory(t *testing.T) { + f := newFakeHost(t) + f.installed(t, linuxInstall("v0.1.0")) + + if err := removeManifest(t.Context(), f.host(nil)); err != nil { + t.Fatalf("removeManifest: %v", err) + } + if _, err := os.Lstat(filepath.Join(f.root, filepath.Dir(ManifestPath))); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("the manifest directory is still there: %v", err) + } +} diff --git a/services/setup/preflight.go b/services/setup/preflight.go new file mode 100644 index 000000000..435742e3f --- /dev/null +++ b/services/setup/preflight.go @@ -0,0 +1,423 @@ +package setup + +import ( + "context" + "errors" + "fmt" + "io/fs" + "net/http" + "os" + "path" + "path/filepath" + "strconv" + "strings" + "syscall" + + "golang.org/x/sys/unix" + + fcapi "github.com/presmihaylov/shard/pkg/firecracker" + "github.com/presmihaylov/shard/services/kernel" + "github.com/presmihaylov/shard/services/sandboxstate" +) + +// finding is what one preflight check found: a failure stops setup before any change, attention is shown and setup goes on. +type finding struct { + check string + lines []string + // attention is a fact the person should know that does not stop setup. + attention bool + // provider marks a failure another provider may not have, so setup offers the provider question again. + provider bool +} + +func failed(lines ...string) *finding { return &finding{lines: lines} } + +func providerFailed(lines ...string) *finding { return &finding{lines: lines, provider: true} } + +func attention(lines ...string) *finding { return &finding{lines: lines, attention: true} } + +type check struct { + title string + run func(ctx context.Context, h Host, l Local) *finding +} + +// The seams a test swaps, since the real ones read the host's mounts or reach the network. +var ( + checkChrootBase = fcapi.CheckChrootBase + kernelURL = kernel.URL + // rootUID owns what a root service runs; a test sets its own uid, since it cannot chown to root. + rootUID = 0 +) + +// The filesystems that clone a disk by reference, by their statfs magic. +const ( + xfsMagic = 0x58465342 + btrfsMagic = 0x9123683e +) + +// firecrackerImageRoom is twice the smallest data image, since datadir sizes the image at half the free space. +const firecrackerImageRoom = 20 << 30 + +// lowDisk is where the disk check warns: one image and a few sandboxes no longer fit. +const lowDisk = 2 << 30 + +// preflight is §8: it checks this host for l, changes nothing, and returns the first failure. +func (s *Setup) preflight(ctx context.Context, l Local) (*finding, error) { + checks := []check{ + {"Supported operating system and CPU", supportedPlatform}, + {"Provider requirements", providerRequirements}, + {"Administrator access", administratorAccess}, + {"Installation paths and permissions", installPaths}, + {"Available disk space and filesystem support", diskSpace}, + {"Download access", downloadAccess}, + {"Existing Shard installation", existingSandboxes}, + } + if l.StartAtBoot { + checks = append(checks, check{"Background service support", serviceSupport}) + } + + titles := make([]string, 0, len(checks)) + for _, c := range checks { + titles = append(titles, c.title) + } + list, err := s.UI.Checklist("Checking this machine", titles) + if err != nil { + return nil, err + } + for i, c := range checks { + if err := list.Start(i); err != nil { + return nil, err + } + f := c.run(ctx, s.Host, l) + switch { + case f == nil: + err = list.Done(i) + case f.attention: + err = list.Attention(i, f.lines...) + default: + f.check = c.title + return f, list.Fail(i, f.lines...) + } + if err != nil { + return nil, err + } + } + + return nil, nil +} + +// supportedPlatform is what a release ships for: shard and shard-init for Linux on x86_64, and the Mac build with VZ. +func supportedPlatform(_ context.Context, h Host, _ Local) *finding { + if h.OS == "linux" && h.Arch == "amd64" || h.OS == "darwin" && h.Arch == "arm64" { + return nil + } + + return failed( + "Shard runs on Linux on x86_64, and on Macs with Apple silicon.", + fmt.Sprintf("This machine runs %s on %s.", osName(h.OS), h.Arch), + ) +} + +func providerRequirements(ctx context.Context, h Host, l Local) *finding { + title := providerTitle(l.Provider) + if lk, ok := lacks(ctx, h, l.Provider); ok { + return providerFailed(title+" requires "+lk.need+".", lk.fact) + } + m := missingTools(h, l.Provider) + if len(m.manual) > 0 { + return providerFailed( + fmt.Sprintf("%s requires %s, which setup cannot install.", title, strings.Join(m.manual, ", ")), + "Install it and run shard setup again.", + ) + } + if _, ok := lookPath(h, "apt-get"); m.apt() && !ok { + return providerFailed( + fmt.Sprintf("%s requires %s, and setup installs them with apt-get, which this machine does not have.", title, strings.Join(m.names(), ", ")), + "Install them and run shard setup again.", + ) + } + if l.Provider != Firecracker { + return nil + } + // Setup keeps a firecracker that is already there, so it must be one the daemon accepts. + binary, ok := lookPath(h, "firecracker") + if !ok { + return nil + } + out, err := h.Run(ctx, rooted(h, binary), "--version") + if err == nil { + err = fcapi.CheckVersionOutput(binary, out) + } + if err != nil { + return providerFailed(sentence(err), "Replace it with Firecracker 1.13.0 or newer, or remove it so setup installs one.") + } + + return nil +} + +func administratorAccess(_ context.Context, h Host, _ Local) *finding { + if h.OS == "darwin" && h.Euid == 0 && h.Env("SUDO_USER") == "" { + return failed( + "On a Mac the daemon runs as your user, and setup cannot tell who that is when it runs as root.", + "Run shard setup as your user. It asks for administrator access when it needs it.", + ) + } + if h.Euid == 0 { + return nil + } + if _, ok := lookPath(h, "sudo"); !ok { + return failed("Setup needs administrator access, and this machine has no sudo.", "Run shard setup as root.") + } + + return nil +} + +// installPaths refuses a place a person other than root could change what the root daemon runs or keeps. +func installPaths(_ context.Context, h Host, _ Local) *finding { + if h.OS == "linux" { + if f := rootOnly(h, binDir); f != nil { + return f + } + } + info, err := os.Lstat(rooted(h, DataDir)) + if errors.Is(err, fs.ErrNotExist) { + return nil + } + if err != nil { + return failed(fmt.Sprintf("Setup could not check %s: %v.", DataDir, err)) + } + if !info.IsDir() { + return failed(DataDir+" exists and is not a directory.", "Move it aside and run shard setup again.") + } + if h.OS != "darwin" { + return nil + } + // On a Mac the daemon runs as the person, and an existing root they do not own refuses its writes. + uid, err := personUID(h) + if err != nil { + return failed(fmt.Sprintf("Setup could not tell which user runs the daemon: %v.", err)) + } + if owner(info) != uid { + return failed( + DataDir+" belongs to another user, and the daemon runs as "+daemonUser(h)+".", + "Change its owner to "+daemonUser(h)+" and run shard setup again.", + ) + } + + return nil +} + +func rootOnly(h Host, dir string) *finding { + info, err := os.Stat(rooted(h, dir)) + if errors.Is(err, fs.ErrNotExist) { + return nil + } + if err != nil { + return failed(fmt.Sprintf("Setup could not check %s: %v.", dir, err)) + } + if !info.IsDir() { + return failed(dir+" exists and is not a directory.", "Move it aside and run shard setup again.") + } + if owner(info) != rootUID || info.Mode().Perm()&0o022 != 0 { + return failed( + dir+" can be changed by a user other than root, and the root daemon runs Shard from it.", + "Make root its owner and remove group and other write access, then run shard setup again.", + ) + } + + return nil +} + +func owner(info fs.FileInfo) int { + st, ok := info.Sys().(*syscall.Stat_t) + if !ok { + return -1 + } + + return int(st.Uid) +} + +// personUID is the uid of daemonUser: the caller, or the one sudo started setup for. +func personUID(h Host) (int, error) { + if h.Euid != 0 { + return h.Euid, nil + } + uid, err := strconv.Atoi(h.Env("SUDO_UID")) + if err != nil { + return 0, fmt.Errorf("read SUDO_UID: %w", err) + } + + return uid, nil +} + +func diskSpace(_ context.Context, h Host, l Local) *finding { + dir, err := nearestDir(rooted(h, DataDir)) + if err != nil { + return failed(fmt.Sprintf("Setup could not check %s: %v.", DataDir, err)) + } + var st unix.Statfs_t + if err := unix.Statfs(dir, &st); err != nil { + return failed(fmt.Sprintf("Setup could not read the filesystem of %s: %v.", DataDir, err)) + } + free := uint64(st.Bavail) * uint64(st.Bsize) //nolint:gosec // G115: a block size is never negative + reflink := int64(st.Type) == xfsMagic || int64(st.Type) == btrfsMagic + if h.OS == "linux" && l.Provider == Firecracker { + if f := firecrackerDisk(h, dir, free, reflink); f != nil { + return f + } + } + if free < lowDisk { + return attention(fmt.Sprintf("Only %s is free for %s.", gib(free), DataDir), "Images and sandboxes may not fit.") + } + + return nil +} + +// firecrackerDisk is datadir's rule ahead of time: a root that cannot clone gets an XFS image, which needs the room and an empty dir. +func firecrackerDisk(h Host, dir string, free uint64, reflink bool) *finding { + // The type stands in for reflink.Probe, which writes a file and so needs root. + if reflink { + if err := checkChrootBase(dir); err != nil { + return providerFailed("Firecracker cannot run its jailer under "+DataDir+".", sentence(err)) + } + + return nil + } + if free < firecrackerImageRoom { + return providerFailed( + "Firecracker needs "+DataDir+" on XFS or Btrfs, or "+gib(firecrackerImageRoom)+" free beside it for an XFS image.", + fmt.Sprintf("%s has %s free, on a filesystem that cannot clone a disk.", strings.TrimPrefix(dir, strings.TrimSuffix(h.Root, "/")), gib(free)), + ) + } + entries, err := os.ReadDir(rooted(h, DataDir)) + switch { + case errors.Is(err, fs.ErrNotExist): + return nil + case errors.Is(err, fs.ErrPermission): + return attention("Setup cannot read "+DataDir+" without administrator access.", "Firecracker mounts an XFS image over it, and the daemon refuses if it holds files.") + case err != nil: + return failed(fmt.Sprintf("Setup could not read %s: %v.", DataDir, err)) + case len(entries) > 0: + return providerFailed( + "Firecracker mounts an XFS image over "+DataDir+", so it must be empty on a filesystem that cannot clone a disk.", + DataDir+" already holds files.", + ) + } + + return nil +} + +// nearestDir is dir, or the closest directory above it that exists. +func nearestDir(dir string) (string, error) { + for { + _, err := os.Stat(dir) + if err == nil { + return dir, nil + } + if !errors.Is(err, fs.ErrNotExist) || dir == filepath.Dir(dir) { + return "", err + } + dir = filepath.Dir(dir) + } +} + +func gib(n uint64) string { return fmt.Sprintf("%.1f GiB", float64(n)/(1<<30)) } + +// downloadAccess reaches every file setup or the first daemon start downloads, before it changes anything. +func downloadAccess(ctx context.Context, h Host, l Local) *finding { + var urls []string + for _, d := range missingTools(h, l.Provider).downloads { + urls = append(urls, d.URL) + } + if h.OS == "linux" { + u, err := AssetURL(ctx, h, h.Version, "shard-init-linux-"+h.Arch) + if err != nil { + return failed(fmt.Sprintf("Setup could not find shard-init for Shard %s: %v.", h.Version, err)) + } + urls = append(urls, u) + } + if l.Provider == Firecracker || l.Provider == VZ { + u, err := kernelURL(h.Arch) + if err != nil { + return failed(sentence(err)) + } + urls = append(urls, u) + } + for _, u := range urls { + if err := reach(ctx, h, u); err != nil { + return failed(fmt.Sprintf("Could not download %s: %v.", path.Base(u), err), "Check the network connection and run shard setup again.") + } + } + + return nil +} + +// reach asks for the first byte of u, which proves the file is there without downloading it. +func reach(ctx context.Context, h Host, u string) (err error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil) + if err != nil { + return err + } + req.Header.Set("Range", "bytes=0-0") + resp, err := h.HTTP.Do(req) + if err != nil { + return err + } + defer func() { err = errors.Join(err, resp.Body.Close()) }() + if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusPartialContent { + return fmt.Errorf("the server answered %s", resp.Status) + } + + return nil +} + +// existingSandboxes refuses a provider other than the one the sandboxes already in the data dir use. +func existingSandboxes(_ context.Context, h Host, l Local) *finding { + recorded, err := sandboxstate.RecordedProvider(rooted(h, DataDir)) + if errors.Is(err, fs.ErrPermission) { + return attention("Setup cannot read "+DataDir+" without administrator access.", "The daemon refuses to start if its sandboxes use another provider.") + } + if err != nil { + return failed(fmt.Sprintf("Setup could not read the sandboxes in %s: %v.", DataDir, err)) + } + if recorded == "" || recorded == l.Provider { + return nil + } + + return providerFailed( + "The sandboxes in "+DataDir+" use "+providerTitle(recorded)+".", + "Setup never changes the provider of existing sandboxes.", + ) +} + +func serviceSupport(_ context.Context, h Host, _ Local) *finding { + manager, tools := "launchd", []string{"launchctl", "plutil"} + if h.OS == "linux" { + manager, tools = "systemd", []string{"systemctl", "systemd-analyze"} + // systemd makes this directory when it boots as PID 1, and only then. + if info, err := os.Stat(rooted(h, "/run/systemd/system")); err != nil || !info.IsDir() { + return noService("Setup configures a systemd service, and systemd does not manage this machine.") + } + } + for _, t := range tools { + if _, ok := lookPath(h, t); !ok { + return noService(fmt.Sprintf("Setup configures a %s service, and this machine has no %s.", manager, t)) + } + } + + return nil +} + +func noService(line string) *finding { + return failed(line, "Run shard setup again and choose No for automatic startup.") +} + +// sentence is an error worded as a line of output. +func sentence(err error) string { + s := err.Error() + if s == "" { + return s + } + + return strings.ToUpper(s[:1]) + s[1:] + "." +} diff --git a/services/setup/preflight_test.go b/services/setup/preflight_test.go new file mode 100644 index 000000000..4e40f20a8 --- /dev/null +++ b/services/setup/preflight_test.go @@ -0,0 +1,315 @@ +package setup + +import ( + "errors" + "os" + "path/filepath" + "slices" + "strings" + "testing" +) + +// preflightOn runs the checks for l on h and returns the finding and the checklist it drew. +func preflightOn(t *testing.T, h Host, l Local) (*finding, *fakeChecklist) { + t.Helper() + ui := &fakeUI{} + f, err := (&Setup{Host: h, UI: ui}).preflight(t.Context(), l) + if err != nil { + t.Fatalf("preflight = %v", err) + } + if len(ui.lists) != 1 { + t.Fatalf("drew %d checklists, want 1", len(ui.lists)) + } + + return f, ui.lists[0] +} + +func wantFinding(t *testing.T, f *finding, check string, provider bool, lines ...string) { + t.Helper() + if f == nil { + t.Fatalf("passed, want %q to fail with %q", check, lines) + } + if f.check != check || f.provider != provider || !slices.Equal(f.lines, lines) { + t.Fatalf("finding %+v, want check %q provider %v lines %q", *f, check, provider, lines) + } +} + +func TestPreflightPassesAReadyHost(t *testing.T) { + l := newLocalHost(t) + + f, list := preflightOn(t, l.host(), Local{Provider: GVisor, StartAtBoot: true}) + if f != nil { + t.Fatalf("a ready host fails %q: %q", f.check, f.lines) + } + want := []string{ + "Supported operating system and CPU", "Provider requirements", "Administrator access", + "Installation paths and permissions", "Available disk space and filesystem support", "Download access", + "Existing Shard installation", "Background service support", + } + if list.title != "Checking this machine" || !slices.Equal(list.steps, want) { + t.Fatalf("checklist %q %q", list.title, list.steps) + } + if slices.ContainsFunc(list.marks, func(m string) bool { return strings.HasPrefix(m, "fail") }) { + t.Fatalf("marks %q", list.marks) + } + if l.changed() { + t.Fatalf("preflight changed the host: %q", l.calls) + } +} + +func TestPreflightStopsAtTheFirstFailure(t *testing.T) { + l := newLocalHost(t) + h := l.host() + h.Arch = "arm64" + + f, list := preflightOn(t, h, Local{Provider: GVisor, StartAtBoot: true}) + + wantFinding(t, f, "Supported operating system and CPU", false, + "Shard runs on Linux on x86_64, and on Macs with Apple silicon.", "This machine runs Linux on arm64.") + want := []string{"start 0", "fail 0: Shard runs on Linux on x86_64, and on Macs with Apple silicon. / This machine runs Linux on arm64."} + if !slices.Equal(list.marks, want) { + t.Fatalf("marks %q, want %q", list.marks, want) + } +} + +func TestPreflightProviderRequirements(t *testing.T) { + cases := []struct { + name string + provider string + mac bool + setup func(l *localHost) + lines []string + }{ + { + name: "no kvm", provider: Firecracker, + lines: []string{"Firecracker requires access to /dev/kvm.", "This machine does not provide it."}, + }, + { + name: "an old firecracker", provider: Firecracker, + setup: func(l *localHost) { + l.write("/dev/kvm", "") + l.write("/usr/local/bin/firecracker", "#!/bin/sh\necho 'Firecracker v1.10.1'\n") + }, + lines: []string{ + "/usr/local/bin/firecracker is firecracker 1.10.1, and shard needs 1.13.0 or newer, whose snapshot load keeps the dirty-page log on.", + "Replace it with Firecracker 1.13.0 or newer, or remove it so setup installs one.", + }, + }, + { + name: "no apt-get for a missing package", provider: GVisor, + setup: func(l *localHost) { + l.remove("/usr/sbin/ip") + l.remove("/usr/bin/apt-get") + }, + lines: []string{ + "gVisor requires ip, and setup installs them with apt-get, which this machine does not have.", + "Install them and run shard setup again.", + }, + }, + { + name: "no codesign on a Mac", provider: VZ, mac: true, + setup: func(l *localHost) { l.remove("/usr/bin/codesign") }, + lines: []string{"macOS Virtualization requires codesign, which setup cannot install.", "Install it and run shard setup again."}, + }, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + l := newLocalHost(t) + h := l.host() + if tc.mac { + h = l.mac() + } + if tc.setup != nil { + tc.setup(l) + } + + f, _ := preflightOn(t, h, Local{Provider: tc.provider}) + + wantFinding(t, f, "Provider requirements", true, tc.lines...) + }) + } +} + +func TestPreflightAdministratorAccess(t *testing.T) { + t.Run("root on a Mac without sudo", func(t *testing.T) { + l := newLocalHost(t) + h := l.mac() + h.Euid = 0 + + f, _ := preflightOn(t, h, Local{Provider: VZ}) + + wantFinding(t, f, "Administrator access", false, + "On a Mac the daemon runs as your user, and setup cannot tell who that is when it runs as root.", + "Run shard setup as your user. It asks for administrator access when it needs it.") + }) + t.Run("a user without sudo", func(t *testing.T) { + l := newLocalHost(t) + l.remove("/usr/bin/sudo") + h := l.host() + h.Euid = 1000 + + f, _ := preflightOn(t, h, Local{Provider: GVisor}) + + wantFinding(t, f, "Administrator access", false, "Setup needs administrator access, and this machine has no sudo.", "Run shard setup as root.") + }) + t.Run("a user with sudo", func(t *testing.T) { + l := newLocalHost(t) + h := l.host() + h.Euid = 1000 + + if f, _ := preflightOn(t, h, Local{Provider: GVisor}); f != nil { + t.Fatalf("fails %q: %q", f.check, f.lines) + } + }) +} + +func TestPreflightInstallPaths(t *testing.T) { + t.Run("a bin dir others can write", func(t *testing.T) { + l := newLocalHost(t) + if err := os.Chmod(filepath.Join(l.root, binDir), 0o777); err != nil { + t.Fatalf("chmod: %v", err) + } + + f, _ := preflightOn(t, l.host(), Local{Provider: GVisor}) + + wantFinding(t, f, "Installation paths and permissions", false, + "/usr/local/bin can be changed by a user other than root, and the root daemon runs Shard from it.", + "Make root its owner and remove group and other write access, then run shard setup again.") + }) + t.Run("a data dir that is a file", func(t *testing.T) { + l := newLocalHost(t) + l.write(DataDir, "") + + f, _ := preflightOn(t, l.host(), Local{Provider: GVisor}) + + wantFinding(t, f, "Installation paths and permissions", false, "/var/lib/shard exists and is not a directory.", "Move it aside and run shard setup again.") + }) + t.Run("a Mac data dir another user owns", func(t *testing.T) { + l := newLocalHost(t) + h := l.mac() + h.Euid = os.Getuid() + 1 + l.mkdir(DataDir) + + f, _ := preflightOn(t, h, Local{Provider: VZ}) + + wantFinding(t, f, "Installation paths and permissions", false, + "/var/lib/shard belongs to another user, and the daemon runs as u.", "Change its owner to u and run shard setup again.") + }) + t.Run("a Mac data dir the person owns", func(t *testing.T) { + l := newLocalHost(t) + h := l.mac() + l.mkdir(DataDir) + swap(t, &kernelURL, func(string) (string, error) { return l.rs.URL + "/download/v0.1.0/shard-init-linux-amd64", nil }) + + if f, _ := preflightOn(t, h, Local{Provider: VZ}); f != nil { + t.Fatalf("fails %q: %q", f.check, f.lines) + } + }) +} + +func TestFirecrackerDisk(t *testing.T) { + const gibs = 1 << 30 + jailer := errors.New("/var/lib/shard is on a filesystem mounted nosuid") + cases := []struct { + name string + reflink bool + chroot error + free uint64 + files bool + want *finding + }{ + {name: "a reflink root the jailer accepts", reflink: true}, + { + name: "a reflink root the jailer refuses", reflink: true, chroot: jailer, + want: providerFailed("Firecracker cannot run its jailer under /var/lib/shard.", "/var/lib/shard is on a filesystem mounted nosuid."), + }, + { + name: "too little room for an image", free: gibs, + want: providerFailed( + "Firecracker needs /var/lib/shard on XFS or Btrfs, or 20.0 GiB free beside it for an XFS image.", + "/var/lib has 1.0 GiB free, on a filesystem that cannot clone a disk.", + ), + }, + { + name: "files where the image mounts", free: 30 * gibs, files: true, + want: providerFailed( + "Firecracker mounts an XFS image over /var/lib/shard, so it must be empty on a filesystem that cannot clone a disk.", + "/var/lib/shard already holds files.", + ), + }, + {name: "room and no data dir", free: 30 * gibs}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + l := newLocalHost(t) + l.mkdir("/var/lib") + if tc.files { + l.write(DataDir+"/state", "") + } + swap(t, &checkChrootBase, func(string) error { return tc.chroot }) + + got := firecrackerDisk(l.host(), filepath.Join(l.root, "/var/lib"), tc.free, tc.reflink) + + if tc.want == nil && got != nil || tc.want != nil && (got == nil || got.provider != tc.want.provider || !slices.Equal(got.lines, tc.want.lines)) { + t.Fatalf("firecrackerDisk = %+v, want %+v", got, tc.want) + } + }) + } +} + +func TestPreflightDownloadAccess(t *testing.T) { + t.Run("no shard-init for this version", func(t *testing.T) { + l := newLocalHost(t) + h := l.host() + h.Version = "v0.2.0" + + f, _ := preflightOn(t, h, Local{Provider: GVisor}) + + if f == nil || f.check != "Download access" || !strings.HasPrefix(f.lines[0], "Setup could not find shard-init for Shard v0.2.0: ") { + t.Fatalf("finding %+v", f) + } + }) + t.Run("a pinned tool the server lacks", func(t *testing.T) { + l := newLocalHost(t) + l.remove("/usr/local/bin/runsc") + swap(t, gvisorRelease, download{Title: "runsc", URL: l.rs.URL + "/download/pins/gvisor.tar.zstd"}) + + f, _ := preflightOn(t, l.host(), Local{Provider: GVisor}) + + wantFinding(t, f, "Download access", false, + "Could not download gvisor.tar.zstd: the server answered 404 Not Found.", "Check the network connection and run shard setup again.") + }) + t.Run("a kernel the server lacks", func(t *testing.T) { + l := newLocalHost(t) + swap(t, &kernelURL, func(arch string) (string, error) { return l.rs.URL + "/download/v0.1.0/vmlinux-" + arch, nil }) + + f, _ := preflightOn(t, l.mac(), Local{Provider: VZ}) + + wantFinding(t, f, "Download access", false, + "Could not download vmlinux-arm64: the server answered 404 Not Found.", "Check the network connection and run shard setup again.") + }) +} + +func TestPreflightExistingSandboxes(t *testing.T) { + l := newLocalHost(t) + l.sandbox(GVisor) + + f, _ := preflightOn(t, l.host(), Local{Provider: Runc}) + wantFinding(t, f, "Existing Shard installation", true, "The sandboxes in /var/lib/shard use gVisor.", "Setup never changes the provider of existing sandboxes.") + + if f, _ := preflightOn(t, l.host(), Local{Provider: GVisor}); f != nil { + t.Fatalf("the recorded provider fails %q: %q", f.check, f.lines) + } +} + +func TestPreflightServiceSupportOnAMac(t *testing.T) { + l := newLocalHost(t) + h := l.mac() + l.remove("/usr/bin/plutil") + swap(t, &kernelURL, func(string) (string, error) { return l.rs.URL + "/download/v0.1.0/shard-init-linux-amd64", nil }) + + f, _ := preflightOn(t, h, Local{Provider: VZ, StartAtBoot: true}) + + wantFinding(t, f, "Background service support", false, + "Setup configures a launchd service, and this machine has no plutil.", "Run shard setup again and choose No for automatic startup.") +} diff --git a/services/setup/providers.go b/services/setup/providers.go new file mode 100644 index 000000000..c41203216 --- /dev/null +++ b/services/setup/providers.go @@ -0,0 +1,194 @@ +package setup + +import ( + "context" + "errors" + "fmt" + "io/fs" + "os" + "path" + "path/filepath" + "strings" + + "github.com/presmihaylov/shard/pkg/vz" +) + +// The names shard daemon --provider takes. +const ( + Firecracker = "firecracker" + GVisor = "gvisor" + Sysbox = "sysbox" + Runc = "runc" + VZ = "vz" +) + +// ProviderChoice is one row of the provider question. +type ProviderChoice struct { + Name string + Title string + Lines []string + // Unavailable says what this host lacks, and is empty when the provider runs here. + Unavailable []string + Recommended bool +} + +type providerText struct { + name, title string + lines []string +} + +const macNeed = "an Apple silicon Mac with macOS 14 or later" + +var providerTexts = []providerText{ + {Firecracker, "Firecracker", []string{ + "Run each sandbox in a small virtual machine.", + "Recommended on Linux when hardware virtualization is available.", + }}, + {GVisor, "gVisor", []string{ + "Run sandboxes with a protective layer between their programs and Linux.", + "Recommended on Linux without hardware virtualization.", + }}, + {Sysbox, "Sysbox", []string{ + "Run Docker and system services inside your sandboxes.", + "Choose this when your workload needs its own Docker environment.", + }}, + {Runc, "runc", []string{ + "Run standard Linux containers that share the host kernel.", + "Choose this for trusted workloads that need standard container behavior.", + }}, + {VZ, "macOS Virtualization", []string{ + "Run each sandbox in a small Linux virtual machine on your Mac.", + "Requires " + macNeed + ".", + }}, +} + +// recommendOrder is the first available provider that wins the recommendation. +var recommendOrder = []string{VZ, Firecracker, GVisor} + +// Providers lists all five providers in the order the question shows them. +func Providers(ctx context.Context, h Host) []ProviderChoice { + rows := make([]ProviderChoice, 0, len(providerTexts)) + for _, p := range providerTexts { + row := ProviderChoice{Name: p.name, Title: p.title, Lines: p.lines} + if l, ok := lacks(ctx, h, p.name); ok { + row.Lines = nil + row.Unavailable = []string{"Requires " + l.need + ".", l.fact} + } + rows = append(rows, row) + } + for _, name := range recommendOrder { + i := providerIndex(name) + if len(rows[i].Unavailable) == 0 { + rows[i].Recommended = true + break + } + } + + return rows +} + +func providerIndex(name string) int { + for i, p := range providerTexts { + if p.name == name { + return i + } + } + + return -1 +} + +// providerTitle is the name a reader sees, and the name itself for one setup does not know. +func providerTitle(name string) string { + i := providerIndex(name) + if i < 0 { + return name + } + + return providerTexts[i].title +} + +// lack is a host capability a provider needs and this host does not have; software setup can install is never one. +type lack struct{ need, fact string } + +func lacks(ctx context.Context, h Host, name string) (lack, bool) { + if name == VZ { + return macLacks(ctx, h) + } + if h.OS != "linux" { + return lack{"a Linux host", "This machine runs " + osName(h.OS) + "."}, true + } + if name != Firecracker { + return lack{}, false + } + + return kvmLacks(h) +} + +func macLacks(ctx context.Context, h Host) (lack, bool) { + if h.OS != "darwin" { + return lack{macNeed, "This machine runs " + osName(h.OS) + "."}, true + } + if h.Arch != "arm64" { + return lack{macNeed, "This Mac has an Intel processor."}, true + } + out, err := h.Run(ctx, "sw_vers", "-productVersion") + if err != nil { + return lack{macNeed, fmt.Sprintf("Setup could not read the macOS version: %v.", err)}, true + } + version := strings.TrimSpace(string(out)) + if !vz.SaveRestoreSupported(h.OS, h.Arch, version) { + return lack{macNeed, "This Mac runs macOS " + version + "."}, true + } + + return lack{}, false +} + +func kvmLacks(h Host) (lack, bool) { + const need = "access to /dev/kvm" + f, err := os.OpenFile(filepath.Join(h.Root, "dev", "kvm"), os.O_RDWR, 0) + if errors.Is(err, fs.ErrNotExist) { + return lack{need, "This machine does not provide it."}, true + } + // Setup runs as the user, and only the root daemon has to open it. + if errors.Is(err, fs.ErrPermission) { + return lack{}, false + } + if pathErr, ok := errors.AsType[*fs.PathError](err); ok { + return lack{need, fmt.Sprintf("/dev/kvm does not open: %v.", pathErr.Err)}, true + } + if err != nil { + return lack{need, fmt.Sprintf("/dev/kvm does not open: %v.", err)}, true + } + if err := f.Close(); err != nil { + return lack{need, fmt.Sprintf("/dev/kvm does not close: %v.", err)}, true + } + + return lack{}, false +} + +func osName(goos string) string { + switch goos { + case "linux": + return "Linux" + case "darwin": + return "macOS" + } + + return goos +} + +// servicePath is where systemd and sudo look for a binary, which is what the daemon finds a runtime on. +var servicePath = []string{"/usr/local/sbin", "/usr/local/bin", "/usr/sbin", "/usr/bin", "/sbin", "/bin"} + +// lookPath finds name where the root daemon looks, never on the user's own PATH. +func lookPath(h Host, name string) (string, bool) { + for _, dir := range servicePath { + p := path.Join(dir, name) + info, err := os.Stat(filepath.Join(h.Root, p)) + if err == nil && info.Mode().IsRegular() && info.Mode()&0o111 != 0 { + return p, true + } + } + + return "", false +} diff --git a/services/setup/release.go b/services/setup/release.go new file mode 100644 index 000000000..19bdaa51f --- /dev/null +++ b/services/setup/release.go @@ -0,0 +1,247 @@ +package setup + +import ( + "bufio" + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "io/fs" + "net/http" + "net/url" + "os" + "path/filepath" + "regexp" + "strconv" + "strings" + + "github.com/presmihaylov/shard/pkg/store" +) + +// sumsAsset is the checksum list every shard release carries beside its binaries. +const sumsAsset = "SHA256SUMS" + +// stableTag is a shard release; SDK and kernel releases share the repository under other tags. +var stableTag = regexp.MustCompile(`^v(\d+)\.(\d+)\.(\d+)$`) + +// Release is one published shard release and the files under it. +type Release struct { + Tag string `json:"tag_name"` + Draft bool `json:"draft"` + Pre bool `json:"prerelease"` + Assets []Asset `json:"assets"` +} + +// Asset is one file of a release. +type Asset struct { + Name string `json:"name"` + URL string `json:"browser_download_url"` +} + +// LatestRelease is the highest stable v<major>.<minor>.<patch> by number, never the Latest flag, so an SDK or kernel release is never picked. +func LatestRelease(ctx context.Context, h Host) (Release, error) { + var releases []Release + if err := getJSON(ctx, h, h.Releases+"?per_page=100", &releases); err != nil { + return Release{}, fmt.Errorf("list shard releases: %w", err) + } + + var latest Release + var latestVersion [3]uint64 + for _, r := range releases { + v, ok := stableVersion(r.Tag) + if r.Draft || r.Pre || !ok { + continue + } + if latest.Tag != "" && !newer(v, latestVersion) { + continue + } + latest, latestVersion = r, v + } + if latest.Tag == "" { + return Release{}, fmt.Errorf("list shard releases: %s has no published v<major>.<minor>.<patch> release", h.Releases) + } + + return latest, nil +} + +// ReleaseByTag is the release named tag, the one a running binary came from. +func ReleaseByTag(ctx context.Context, h Host, tag string) (Release, error) { + var r Release + if err := getJSON(ctx, h, h.Releases+"/tags/"+url.PathEscape(tag), &r); err != nil { + return Release{}, fmt.Errorf("find shard release %s: %w", tag, err) + } + + return r, nil +} + +// FetchAsset downloads the file name of release tag to dst, and only once its hash matches the release's SHA256SUMS. +func FetchAsset(ctx context.Context, h Host, tag, name, dst string, perm fs.FileMode) error { + r, err := ReleaseByTag(ctx, h, tag) + if err != nil { + return err + } + + return r.Fetch(ctx, h, name, dst, perm) +} + +// AssetURL is where release tag serves the file name, for a check that probes access without a download. +func AssetURL(ctx context.Context, h Host, tag, name string) (string, error) { + r, err := ReleaseByTag(ctx, h, tag) + if err != nil { + return "", err + } + a, err := r.asset(name) + if err != nil { + return "", err + } + + return a.URL, nil +} + +// Fetch downloads the file name of r to dst through a sibling part file, renamed into place after the hash matched. +func (r Release) Fetch(ctx context.Context, h Host, name, dst string, perm fs.FileMode) error { + want, err := r.sum(ctx, h, name) + if err != nil { + return err + } + a, err := r.asset(name) + if err != nil { + return err + } + + resp, err := get(ctx, h, a.URL, "") + if err != nil { + return fmt.Errorf("download %s from shard release %s: %w", name, r.Tag, err) + } + defer resp.Body.Close() + + f, err := os.CreateTemp(filepath.Dir(dst), filepath.Base(dst)+".*.part") + if err != nil { + return fmt.Errorf("download %s: %w", name, err) + } + part := f.Name() + hash := sha256.New() + _, err = io.Copy(io.MultiWriter(f, hash), resp.Body) + if err == nil { + err = f.Chmod(perm) + } + if err == nil { + err = f.Sync() + } + if err := errors.Join(err, f.Close()); err != nil { + return errors.Join(fmt.Errorf("download %s from shard release %s: %w", name, r.Tag, err), os.Remove(part)) + } + + if got := hex.EncodeToString(hash.Sum(nil)); got != want { + return errors.Join(fmt.Errorf("verify %s from shard release %s: sha256 is %s, SHA256SUMS says %s", name, r.Tag, got, want), os.Remove(part)) + } + if err := os.Rename(part, dst); err != nil { + return errors.Join(fmt.Errorf("download %s: %w", name, err), os.Remove(part)) + } + + return store.SyncDir(filepath.Dir(dst)) +} + +// sum is the hash SHA256SUMS of r gives for name. +func (r Release) sum(ctx context.Context, h Host, name string) (string, error) { + a, err := r.asset(sumsAsset) + if err != nil { + return "", err + } + + resp, err := get(ctx, h, a.URL, "") + if err != nil { + return "", fmt.Errorf("download %s of shard release %s: %w", sumsAsset, r.Tag, err) + } + defer resp.Body.Close() + + // sha256sum writes "<hash> <name>", or "<hash> *<name>" in binary mode. + lines := bufio.NewScanner(io.LimitReader(resp.Body, 1<<20)) + for lines.Scan() { + sum, file, ok := strings.Cut(lines.Text(), " ") + if ok && strings.TrimLeft(file, " *") == name { + return strings.ToLower(sum), nil + } + } + if err := lines.Err(); err != nil { + return "", fmt.Errorf("read %s of shard release %s: %w", sumsAsset, r.Tag, err) + } + + return "", fmt.Errorf("%s of shard release %s has no line for %s", sumsAsset, r.Tag, name) +} + +func (r Release) asset(name string) (Asset, error) { + for _, a := range r.Assets { + if a.Name == name { + return a, nil + } + } + + return Asset{}, fmt.Errorf("shard release %s has no file %s", r.Tag, name) +} + +func stableVersion(tag string) ([3]uint64, bool) { + m := stableTag.FindStringSubmatch(tag) + if m == nil { + return [3]uint64{}, false + } + + var v [3]uint64 + for i := range v { + n, err := strconv.ParseUint(m[i+1], 10, 64) + if err != nil { + return [3]uint64{}, false + } + v[i] = n + } + + return v, true +} + +func newer(a, b [3]uint64) bool { + for i := range a { + if a[i] != b[i] { + return a[i] > b[i] + } + } + + return false +} + +func getJSON(ctx context.Context, h Host, u string, v any) error { + resp, err := get(ctx, h, u, "application/vnd.github+json") + if err != nil { + return err + } + defer resp.Body.Close() + + if err := json.NewDecoder(resp.Body).Decode(v); err != nil { + return fmt.Errorf("decode %s: %w", u, err) + } + + return nil +} + +// get answers only a 200, so a rate limit or a missing tag names its status and never reads as an empty list. +func get(ctx context.Context, h Host, u, accept string) (*http.Response, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil) + if err != nil { + return nil, err + } + if accept != "" { + req.Header.Set("Accept", accept) + } + + resp, err := h.HTTP.Do(req) + if err != nil { + return nil, err + } + if resp.StatusCode != http.StatusOK { + return nil, errors.Join(fmt.Errorf("GET %s: %s", u, resp.Status), resp.Body.Close()) + } + + return resp, nil +} diff --git a/services/setup/release_test.go b/services/setup/release_test.go new file mode 100644 index 000000000..c023de13b --- /dev/null +++ b/services/setup/release_test.go @@ -0,0 +1,219 @@ +package setup + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "sync/atomic" + "testing" +) + +// releaseServer is a fake releases API whose assets point back at itself. +type releaseServer struct { + *httptest.Server + releases []Release + files map[string]string + downloads atomic.Int32 +} + +func newReleaseServer(t *testing.T) *releaseServer { + t.Helper() + + rs := &releaseServer{files: map[string]string{}} + mux := http.NewServeMux() + mux.HandleFunc("GET /releases", func(w http.ResponseWriter, r *http.Request) { + if r.URL.Query().Get("per_page") != "100" { + http.Error(w, "want per_page=100", http.StatusBadRequest) + return + } + writeJSON(t, w, rs.releases) + }) + mux.HandleFunc("GET /releases/tags/{tag}", func(w http.ResponseWriter, r *http.Request) { + for _, rel := range rs.releases { + if rel.Tag == r.PathValue("tag") { + writeJSON(t, w, rel) + return + } + } + http.NotFound(w, r) + }) + mux.HandleFunc("GET /download/{tag}/{name}", func(w http.ResponseWriter, r *http.Request) { + rs.downloads.Add(1) + body, ok := rs.files[r.PathValue("tag")+"/"+r.PathValue("name")] + if !ok { + http.NotFound(w, r) + return + } + if _, err := w.Write([]byte(body)); err != nil { + t.Errorf("write %s: %v", r.URL.Path, err) + } + }) + rs.Server = httptest.NewServer(mux) + t.Cleanup(rs.Close) + + return rs +} + +// add publishes a release with files, and a SHA256SUMS for them unless sums is false. +func (rs *releaseServer) add(tag string, draft, pre, sums bool, files map[string]string) { + rel := Release{Tag: tag, Draft: draft, Pre: pre} + var list strings.Builder + for name, body := range files { + rs.files[tag+"/"+name] = body + rel.Assets = append(rel.Assets, Asset{Name: name, URL: rs.URL + "/download/" + tag + "/" + name}) + sum := sha256.Sum256([]byte(body)) + list.WriteString(hex.EncodeToString(sum[:]) + " " + name + "\n") + } + if sums { + rs.files[tag+"/"+sumsAsset] = list.String() + rel.Assets = append(rel.Assets, Asset{Name: sumsAsset, URL: rs.URL + "/download/" + tag + "/" + sumsAsset}) + } + rs.releases = append(rs.releases, rel) +} + +func (rs *releaseServer) host() Host { + return Host{Releases: rs.URL + "/releases", HTTP: rs.Client()} +} + +func writeJSON(t *testing.T, w http.ResponseWriter, v any) { + t.Helper() + if err := json.NewEncoder(w).Encode(v); err != nil { + t.Errorf("encode: %v", err) + } +} + +func TestLatestReleasePicksTheHighestStableTag(t *testing.T) { + rs := newReleaseServer(t) + // API order is newest first; the pick must not trust it, nor compare tags as text. + rs.add("sdk-typescript-v0.3.0", false, false, true, nil) + rs.add("v0.9.1", false, false, true, nil) + rs.add("kernel-6.12.110-3", false, false, true, nil) + rs.add("v0.11.0-rc.1", false, false, true, nil) + rs.add("v0.12.0", false, true, true, nil) + rs.add("v0.13.0", true, false, true, nil) + rs.add("v0.10.0", false, false, true, nil) + rs.add("sdk-python-v0.2.0", false, false, true, nil) + + got, err := LatestRelease(context.Background(), rs.host()) + if err != nil { + t.Fatalf("LatestRelease: %v", err) + } + if got.Tag != "v0.10.0" { + t.Fatalf("LatestRelease picked %s, want v0.10.0", got.Tag) + } +} + +func TestLatestReleaseRefusesAListWithNoStableTag(t *testing.T) { + rs := newReleaseServer(t) + rs.add("sdk-typescript-v0.3.0", false, false, true, nil) + rs.add("v1.0.0", false, true, true, nil) + + _, err := LatestRelease(context.Background(), rs.host()) + if err == nil || !strings.Contains(err.Error(), "no published v<major>.<minor>.<patch> release") { + t.Fatalf("LatestRelease error = %v, want the missing stable release named", err) + } +} + +func TestLatestReleaseNamesAnAPIFailure(t *testing.T) { + api := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + http.Error(w, "rate limited", http.StatusForbidden) + })) + t.Cleanup(api.Close) + + _, err := LatestRelease(context.Background(), Host{Releases: api.URL + "/releases", HTTP: api.Client()}) + if err == nil || !strings.Contains(err.Error(), "403 Forbidden") { + t.Fatalf("LatestRelease error = %v, want the 403 named", err) + } +} + +func TestFetchAssetWritesOnlyAVerifiedFile(t *testing.T) { + rs := newReleaseServer(t) + rs.add("v0.2.0", false, false, true, map[string]string{"shard-linux-amd64": "new binary"}) + dst := filepath.Join(t.TempDir(), "shard") + + if err := FetchAsset(context.Background(), rs.host(), "v0.2.0", "shard-linux-amd64", dst, 0o755); err != nil { + t.Fatalf("FetchAsset: %v", err) + } + data, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("read %s: %v", dst, err) + } + if string(data) != "new binary" { + t.Fatalf("FetchAsset wrote %q", data) + } + info, err := os.Stat(dst) + if err != nil { + t.Fatalf("stat %s: %v", dst, err) + } + if info.Mode().Perm() != 0o755 { + t.Fatalf("FetchAsset mode = %v, want 0755", info.Mode().Perm()) + } +} + +func TestFetchAssetRefusesABadOrMissingSum(t *testing.T) { + cases := map[string]struct { + asset string + sums bool + tamper bool + want string + }{ + "no SHA256SUMS": {asset: "shard-linux-amd64", sums: false, want: "has no file SHA256SUMS"}, + "hash mismatch": {asset: "shard-linux-amd64", sums: true, tamper: true, want: "SHA256SUMS says"}, + "no line for it": {asset: "shard-init-linux-amd64", sums: true, want: "has no line for shard-init-linux-amd64"}, + } + for name, tc := range cases { + t.Run(name, func(t *testing.T) { + rs := newReleaseServer(t) + rs.add("v0.2.0", false, false, tc.sums, map[string]string{"shard-linux-amd64": "new binary"}) + if tc.tamper { + rs.files["v0.2.0/shard-linux-amd64"] = "tampered" + } + dir := t.TempDir() + dst := filepath.Join(dir, "shard") + + err := FetchAsset(context.Background(), rs.host(), "v0.2.0", tc.asset, dst, 0o755) + if err == nil || !strings.Contains(err.Error(), tc.want) { + t.Fatalf("FetchAsset error = %v, want %q", err, tc.want) + } + entries, err := os.ReadDir(dir) + if err != nil { + t.Fatalf("read %s: %v", dir, err) + } + if len(entries) != 0 { + t.Fatalf("FetchAsset left %v behind", entries) + } + }) + } +} + +func TestFetchAssetNamesAMissingTag(t *testing.T) { + rs := newReleaseServer(t) + rs.add("v0.2.0", false, false, true, nil) + + err := FetchAsset(context.Background(), rs.host(), "dev-abc123", "shard-init-linux-amd64", filepath.Join(t.TempDir(), "x"), 0o755) + if err == nil || !strings.Contains(err.Error(), "find shard release dev-abc123") || !strings.Contains(err.Error(), "404") { + t.Fatalf("FetchAsset error = %v, want the missing tag and 404 named", err) + } +} + +func TestAssetURLNamesAMissingFile(t *testing.T) { + rs := newReleaseServer(t) + rs.add("v0.2.0", false, false, true, map[string]string{"shard-init-linux-amd64": "init"}) + + got, err := AssetURL(context.Background(), rs.host(), "v0.2.0", "shard-init-linux-amd64") + if err != nil || got != rs.URL+"/download/v0.2.0/shard-init-linux-amd64" { + t.Fatalf("AssetURL = %q, %v", got, err) + } + if rs.downloads.Load() != 0 { + t.Fatalf("AssetURL downloaded %d files", rs.downloads.Load()) + } + if _, err := AssetURL(context.Background(), rs.host(), "v0.2.0", "shard-init-linux-arm64"); err == nil || !strings.Contains(err.Error(), "shard-init-linux-arm64") { + t.Fatalf("AssetURL error = %v, want the missing file named", err) + } +} diff --git a/services/setup/remote.go b/services/setup/remote.go new file mode 100644 index 000000000..c9f719470 --- /dev/null +++ b/services/setup/remote.go @@ -0,0 +1,340 @@ +package setup + +import ( + "context" + "crypto/tls" + "errors" + "fmt" + "net" + "net/http" + "net/url" + "path/filepath" + "strings" + "syscall" + + "github.com/presmihaylov/shard/pkg/term" + "github.com/presmihaylov/shard/services/client" +) + +// verifySteps are the §12 checks, in the order Verify runs them. +var verifySteps = []string{"Reach the server", "Verify authentication", "Read server capabilities"} + +// The choices after a failed check, in the order the wizard shows them. +const ( + retryAgain = iota + retryEdit + retryExit +) + +var retryOptions = []term.Option{ + retryAgain: {Name: "retry", Label: "Retry", Default: true}, + retryEdit: {Name: "edit", Label: "Edit connection details"}, + retryExit: {Name: "exit", Label: "Exit"}, +} + +// remote is §12 and §13: the saved connection's menu when there is one, else a new connection. +func (s *Setup) remote(ctx context.Context) error { + path, err := client.ConfigPath(s.Host.Env) + if err != nil { + return err + } + saved, err := client.LoadConfig(path) + if err != nil { + return err + } + if saved.Remote != "" { + return s.saved(ctx, path, saved) + } + + return s.connect(ctx, path, client.Config{}) +} + +// connect asks for a connection, verifies it, and offers to save it; previous stays on disk until the new one is written. +func (s *Setup) connect(ctx context.Context, path string, previous client.Config) error { + conn, err := s.ask(ctx, true) + if err != nil { + return err + } + conn, caps, err := s.verify(ctx, conn) + if err != nil { + return err + } + + return s.offerSave(ctx, path, previous, conn, caps) +} + +// verify checks conn until it passes or the user exits, and returns the connection that passed, edited or not. +func (s *Setup) verify(ctx context.Context, conn client.Config) (client.Config, client.Capabilities, error) { + for { + list, err := s.UI.Checklist("Checking the connection", verifySteps) + if err != nil { + return client.Config{}, client.Capabilities{}, err + } + caps, verifyErr := Verify(ctx, s.Host, conn, list) + if verifyErr == nil { + return conn, caps, nil + } + + choice, err := s.UI.Select(ctx, AskRetry, "What would you like to do?", retryOptions) + if err != nil { + return client.Config{}, client.Capabilities{}, errors.Join(verifyErr, err) + } + if choice == retryExit { + return client.Config{}, client.Capabilities{}, verifyErr + } + if choice == retryEdit { + if conn, err = s.ask(ctx, false); err != nil { + return client.Config{}, client.Capabilities{}, err + } + } + } +} + +// ask reads the URL and the key; with envKey a key in SHARD_API_KEY is used rather than asked for, and an edit asks the person. +func (s *Setup) ask(ctx context.Context, envKey bool) (client.Config, error) { + remote, err := s.UI.Text(ctx, AskURL, "Shard server URL:") + if err != nil { + return client.Config{}, err + } + remote = strings.TrimSpace(remote) + + if parsed, err := url.Parse(remote); err == nil && parsed.Scheme == "http" { + if err := s.UI.Print("", "HTTP does not encrypt your API key or requests.", "Use it only through a trusted encrypted network.", ""); err != nil { + return client.Config{}, err + } + proceed, err := s.UI.Confirm(ctx, AskHTTP, "Continue?", false) + if err != nil { + return client.Config{}, err + } + if !proceed { + return client.Config{}, ErrDeclined + } + } + + if key := strings.TrimSpace(s.Host.Env(client.APIKeyEnv)); envKey && key != "" { + if err := s.UI.Print("Using the API key in " + client.APIKeyEnv + "."); err != nil { + return client.Config{}, err + } + + return client.Config{Remote: remote, APIKey: key}, nil + } + + key, err := s.UI.Secret(ctx, AskAPIKey, "API key:") + if err != nil { + return client.Config{}, err + } + key = strings.TrimSpace(key) + if key == "" { + return client.Config{}, fmt.Errorf("setup needs an API key: enter one at the prompt, or set %s", client.APIKeyEnv) + } + + return client.Config{Remote: remote, APIKey: key}, nil +} + +// Verify runs the §12 checks of conn on list and returns the server's capabilities; a failed step shows what to do about it, so its error is a StoppedError. +func Verify(ctx context.Context, host Host, conn client.Config, list Checklist) (client.Capabilities, error) { + var c *client.Client + var caps client.Capabilities + checks := []func() ([]string, error){ + func() ([]string, error) { + ca, err := client.ReadCA(conn.Remote, host.Env(client.CAFileEnv)) + if err != nil { + return []string{err.Error()}, err + } + if c, err = client.NewRemote(conn.Remote, conn.APIKey, ca); err != nil { + return []string{err.Error()}, err + } + if err := c.Reach(ctx); err != nil { + return reachDetail(conn.Remote, err), err + } + + return nil, nil + }, + func() ([]string, error) { + if _, err := c.Version(ctx); err != nil { + return authDetail(err), fmt.Errorf("verify the API key with %s: %w", client.Redacted(conn.Remote), err) + } + + return nil, nil + }, + func() (detail []string, err error) { + if caps, err = c.Capabilities(ctx); err != nil { + return []string{err.Error()}, fmt.Errorf("read the capabilities of %s: %w", client.Redacted(conn.Remote), err) + } + + return nil, nil + }, + } + + for i, check := range checks { + if err := list.Start(i); err != nil { + return client.Capabilities{}, err + } + if detail, err := check(); err != nil { + return client.Capabilities{}, errors.Join(&StoppedError{Step: verifySteps[i], Err: err}, list.Fail(i, detail...)) + } + if err := list.Done(i); err != nil { + return client.Capabilities{}, err + } + } + + return caps, nil +} + +// reachDetail says why remote did not answer, and explains SHARD_CA_FILE for a certificate this machine does not trust; there is no way to skip the check. +func reachDetail(remote string, err error) []string { + var untrusted *tls.CertificateVerificationError + if errors.As(err, &untrusted) { + return []string{ + "This machine does not trust the server certificate.", + "For a private certificate authority, exit, set " + client.CAFileEnv + " to its PEM file, and run shard setup again.", + } + } + var unreachable *client.ConnectError + if errors.As(err, &unreachable) { + return []string{ + "Could not reach " + client.Redacted(remote) + ": " + dialCause(unreachable.Err) + ".", + "Check the URL, and that shard serve or the proxy in front of it runs.", + } + } + + return []string{err.Error()} +} + +// dialCause words the common dial failures as a person reads them, and leaves any other as the dialer said it. +func dialCause(err error) string { + var dns *net.DNSError + var timeout net.Error + switch { + case errors.As(err, &dns) && dns.IsNotFound: + return "no such host" + case errors.Is(err, syscall.ECONNREFUSED): + return "connection refused" + case errors.Is(err, context.DeadlineExceeded), errors.As(err, &timeout) && timeout.Timeout(): + return "connection timed out" + } + + return err.Error() +} + +// authDetail words the 401 of a server that refused the key. +func authDetail(err error) []string { + var refused *client.APIError + if errors.As(err, &refused) && refused.Status == http.StatusUnauthorized { + return []string{"The server did not accept the API key.", "Check the key, or ask the server administrator for a new one."} + } + + return []string{err.Error()} +} + +// offerSave shows the verified connection and saves it on the user's word. +func (s *Setup) offerSave(ctx context.Context, path string, previous, conn client.Config, caps client.Capabilities) error { + lines := append(connectedLines(conn, caps), "Saving stores your API key as plain text in a file only your user can read.", "") + if err := s.UI.Print(lines...); err != nil { + return err + } + + save, err := s.UI.Confirm(ctx, AskSave, "Save this connection for future Shard commands?", true) + if err != nil { + return err + } + if !save { + return s.UI.Print(notSavedLines(previous, conn, s.Host.Env)...) + } + + if err := client.SaveConfig(path, conn); err != nil { + return fmt.Errorf("save the connection: %w", err) + } + abs, err := filepath.Abs(path) + if err != nil { + return fmt.Errorf("resolve %s: %w", path, err) + } + + return s.UI.Print(savedLines(abs, conn, s.Host.Env)...) +} + +// connectedLines are the §13 result lines, then every lifecycle capability, the unsupported ones too. +func connectedLines(conn client.Config, caps client.Capabilities) []string { + lines := []string{ + "✓ Connected to " + client.Redacted(conn.Remote), + "✓ API key accepted", + "✓ Server capabilities retrieved", + "", + "Server capabilities:", + } + for _, capability := range []struct { + verb string + supported bool + }{ + {"create", caps.Create}, {"start", caps.Start}, {"stop", caps.Stop}, {"remove", caps.Remove}, + {"pause", caps.Pause}, {"resume", caps.Resume}, {"fork", caps.Fork}, {"snapshot", caps.Snapshot}, + } { + support := "supported" + if !capability.supported { + support = "not supported" + } + lines = append(lines, fmt.Sprintf(" %-10s %s", capability.verb, support)) + } + + return append(lines, "", "Capabilities show what the server supports. They do not override the permissions of your API key.", "") +} + +// savedLines are the §13 completion text, with the variables that still beat what the file says. +func savedLines(path string, conn client.Config, env func(string) string) []string { + lines := []string{ + "✓ Connection saved", + "", + "Configuration: " + path, + "", + "The file contains your API key and is accessible only to your user.", + "Shard will use this connection automatically.", + } + if remote := strings.TrimSpace(env(client.RemoteEnv)); remote != "" { + lines = append(lines, "", client.RemoteEnv+" is set to "+client.Redacted(remote)+", and Shard commands use it before the saved connection.") + } + if key := strings.TrimSpace(env(client.APIKeyEnv)); key != "" && key != conn.APIKey { + lines = append(lines, "", client.APIKeyEnv+" is set to another key, and Shard commands use it before the saved one.") + } + if caFile := env(client.CAFileEnv); caFile != "" { + lines = append(lines, "", "Keep "+client.CAFileEnv+" set: the saved connection does not store the certificate authority.") + } + + return append(lines, + "", + "Next steps:", + "", + " List sandboxes:", + " shard list", + "", + " Create a sandbox:", + " shard create --name demo --memory 512MiB alpine:3.20", + "", + " Run a command:", + " shard exec demo echo hello", + "", + " Remove the sandbox:", + " shard remove --force demo", + "", + "Documentation: https://useshards.com/docs", + ) +} + +// notSavedLines say how to use the verified connection without the file; the key stays a placeholder. +func notSavedLines(previous, conn client.Config, env func(string) string) []string { + lines := []string{"The connection was verified but not saved."} + if previous.Remote != "" { + lines = append(lines, "The saved connection to "+client.Redacted(previous.Remote)+" is unchanged.") + } + lines = append(lines, + "", + "To use it, set these environment variables:", + "", + " export "+client.RemoteEnv+"="+client.Redacted(conn.Remote), + " export "+client.APIKeyEnv+"=<your API key>", + ) + if caFile := env(client.CAFileEnv); caFile != "" { + lines = append(lines, " export "+client.CAFileEnv+"="+caFile) + } + + return lines +} diff --git a/services/setup/remote_test.go b/services/setup/remote_test.go new file mode 100644 index 000000000..bff6bb4d7 --- /dev/null +++ b/services/setup/remote_test.go @@ -0,0 +1,692 @@ +package setup + +import ( + "context" + "crypto/tls" + "encoding/pem" + "errors" + "fmt" + "io" + "log" + "maps" + "net" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "slices" + "strings" + "sync/atomic" + "syscall" + "testing" + + "github.com/presmihaylov/shard/services/client" +) + +// testKey is a synthetic key no real credential holds, so any line that carries it is a leak. +const testKey = "shard657-synthetic-key-9e8d7c6b" + +// someCaps leaves fork and snapshot unsupported, so a test sees both words. +const someCaps = `{"create":true,"start":true,"stop":true,"remove":true,"pause":true,"resume":true,"fork":false,"snapshot":false}` + +// front is a shard serve that takes key alone; fail answers a request with a 503 while it returns true. +func front(key string, fail func() bool) *httptest.Server { + server := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + body, status := `{"version":"test"}`, http.StatusOK + switch { + case r.Header.Get("Authorization") != "Bearer "+key: + body, status = `{"error":{"code":"unauthorized","message":"the bearer token is missing or invalid"}}`, http.StatusUnauthorized + case fail != nil && fail(): + body, status = `{"error":{"code":"unavailable","message":"the daemon is restarting"}}`, http.StatusServiceUnavailable + case r.URL.Path == "/v0/capabilities": + body = someCaps + } + w.WriteHeader(status) + if _, err := w.Write([]byte(body)); err != nil { + panic(fmt.Sprintf("write the response: %v", err)) + } + })) + // The untrusted certificate test fails a handshake on purpose. + server.Config.ErrorLog = log.New(io.Discard, "", 0) + + return server +} + +// tlsFront is front over https, with the CA file that trusts it. +func tlsFront(t *testing.T, key string, fail func() bool) (string, string) { + t.Helper() + + server := front(key, fail) + server.StartTLS() + t.Cleanup(server.Close) + + ca := filepath.Join(t.TempDir(), "ca.pem") + if err := os.WriteFile(ca, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: server.Certificate().Raw}), 0o600); err != nil { + t.Fatalf("write the ca file: %v", err) + } + + return server.URL, ca +} + +// testHost is a machine whose environment is vars alone, with its configuration directory in a temp dir. +func testHost(t *testing.T, vars map[string]string) (Host, string) { + t.Helper() + + env := map[string]string{client.ConfigHomeEnv: t.TempDir()} + maps.Copy(env, vars) + host := Host{Env: func(name string) string { return env[name] }} + path, err := client.ConfigPath(host.Env) + if err != nil { + t.Fatalf("ConfigPath: %v", err) + } + + return host, path +} + +// saveConnection writes c where the wizard finds it, as an earlier run would have. +func saveConnection(t *testing.T, path string, c client.Config) { + t.Helper() + + if err := client.SaveConfig(path, c); err != nil { + t.Fatalf("save the connection: %v", err) + } +} + +// savedConnection is what the file at path holds now, the zero Config for no file. +func savedConnection(t *testing.T, path string) client.Config { + t.Helper() + + c, err := client.LoadConfig(path) + if err != nil { + t.Fatalf("load the connection: %v", err) + } + + return c +} + +// noLeak fails a test whose output or error carries the key. +func noLeak(t *testing.T, ui *fakeUI, err error) { + t.Helper() + + for _, line := range ui.printed { + if strings.Contains(line, testKey) { + t.Errorf("printed the key: %q", line) + } + } + for _, list := range ui.lists { + for _, mark := range list.marks { + if strings.Contains(mark, testKey) { + t.Errorf("the checklist shows the key: %q", mark) + } + } + } + if err != nil && strings.Contains(err.Error(), testKey) { + t.Errorf("the error quotes the key: %v", err) + } +} + +// names are the Name of each option a question offered, in order. +func names(ui *fakeUI, q Question) []string { + var out []string + for _, o := range ui.options[q] { + out = append(out, o.Name) + } + + return out +} + +func containsAll(t *testing.T, printed []string, want ...string) { + t.Helper() + + for _, line := range want { + if !slices.Contains(printed, line) { + t.Errorf("printed no line %q in %q", line, printed) + } + } +} + +var allDone = []string{"start 0", "done 0", "start 1", "done 1", "start 2", "done 2"} + +// The §13 completion text, quoted from the spec, after the path line. +var completion = []string{ + "The file contains your API key and is accessible only to your user.", + "Shard will use this connection automatically.", + "", + "Next steps:", + "", + " List sandboxes:", + " shard list", + "", + " Create a sandbox:", + " shard create --name demo --memory 512MiB alpine:3.20", + "", + " Run a command:", + " shard exec demo echo hello", + "", + " Remove the sandbox:", + " shard remove --force demo", + "", + "Documentation: https://useshards.com/docs", +} + +// A connection is checked in three steps, shows every capability, and is saved 0600 on a yes. (SHARD-657) +func TestRemoteVerifiesThenSavesTheConnection(t *testing.T) { + url, ca := tlsFront(t, testKey, nil) + host, path := testHost(t, map[string]string{client.CAFileEnv: ca}) + ui := &fakeUI{ + selects: map[Question]string{AskMode: "remote"}, + texts: map[Question]string{AskURL: url}, + secrets: map[Question]string{AskAPIKey: testKey}, + confirms: map[Question]bool{AskSave: true}, + } + + err := (&Setup{Host: host, UI: ui}).Run(t.Context()) + if err != nil { + t.Fatalf("Run: %v", err) + } + + if want := []Question{AskMode, AskURL, AskAPIKey, AskSave}; !slices.Equal(ui.asked, want) { + t.Errorf("asked %v, want %v", ui.asked, want) + } + if len(ui.lists) != 1 || ui.lists[0].title != "Checking the connection" || !slices.Equal(ui.lists[0].steps, verifySteps) { + t.Fatalf("drew %d checklists, want the one of %v", len(ui.lists), verifySteps) + } + if !slices.Equal(ui.lists[0].marks, allDone) { + t.Errorf("marked %v, want %v", ui.lists[0].marks, allDone) + } + if got := savedConnection(t, path); got != (client.Config{Remote: url, APIKey: testKey}) { + t.Errorf("saved %v, want the verified connection", got) + } + info, err := os.Stat(path) + if err != nil { + t.Fatalf("stat the configuration: %v", err) + } + if info.Mode().Perm() != 0o600 { + t.Errorf("the configuration has mode %o, want 600", info.Mode().Perm()) + } + + containsAll(t, ui.printed, + "✓ Connected to "+url, "✓ API key accepted", "✓ Server capabilities retrieved", + " create supported", " resume supported", " fork not supported", " snapshot not supported", + "Capabilities show what the server supports. They do not override the permissions of your API key.", + "Saving stores your API key as plain text in a file only your user can read.", + "✓ Connection saved", "Configuration: "+path, + "Keep "+client.CAFileEnv+" set: the saved connection does not store the certificate authority.", + ) + for _, block := range [][]string{completion[:2], completion[2:]} { + if !strings.Contains(strings.Join(ui.printed, "\n"), strings.Join(block, "\n")) { + t.Errorf("printed %q, want the §13 completion text %q", ui.printed, block) + } + } + noLeak(t, ui, err) +} + +// HTTP warns and asks first; a no stops before anything is sent, and a yes connects. (SHARD-657) +func TestAnHTTPServerAsksFirst(t *testing.T) { + server := front(testKey, nil) + server.Start() + t.Cleanup(server.Close) + host, path := testHost(t, nil) + + ui := &fakeUI{texts: map[Question]string{AskURL: server.URL}, confirms: map[Question]bool{AskHTTP: false}} + err := (&Setup{Host: host, UI: ui}).remote(t.Context()) + if !errors.Is(err, ErrDeclined) { + t.Fatalf("remote after a no: %v, want ErrDeclined", err) + } + if want := []Question{AskURL, AskHTTP}; !slices.Equal(ui.asked, want) { + t.Errorf("asked %v, want %v", ui.asked, want) + } + containsAll(t, ui.printed, "HTTP does not encrypt your API key or requests.", "Use it only through a trusted encrypted network.") + if len(ui.lists) != 0 { + t.Errorf("checked the connection after a no") + } + + ui = &fakeUI{ + texts: map[Question]string{AskURL: server.URL}, + secrets: map[Question]string{AskAPIKey: testKey}, + confirms: map[Question]bool{AskHTTP: true, AskSave: true}, + } + if err := (&Setup{Host: host, UI: ui}).remote(t.Context()); err != nil { + t.Fatalf("remote after a yes: %v", err) + } + if got := savedConnection(t, path); got.Remote != server.URL { + t.Errorf("saved %v, want %s", got, server.URL) + } +} + +// SHARD_API_KEY answers the key, so a person or a script is not asked for it again. (SHARD-657) +func TestSHARDAPIKEYIsUsedRatherThanAsked(t *testing.T) { + url, ca := tlsFront(t, testKey, nil) + host, path := testHost(t, map[string]string{client.CAFileEnv: ca, client.APIKeyEnv: testKey}) + ui := &fakeUI{texts: map[Question]string{AskURL: url}, confirms: map[Question]bool{AskSave: true}} + + err := (&Setup{Host: host, UI: ui}).remote(t.Context()) + if err != nil { + t.Fatalf("remote: %v", err) + } + + if slices.Contains(ui.asked, AskAPIKey) { + t.Errorf("asked for the key with %s set", client.APIKeyEnv) + } + containsAll(t, ui.printed, "Using the API key in "+client.APIKeyEnv+".") + if got := savedConnection(t, path); got.APIKey != testKey { + t.Errorf("saved a different key than %s", client.APIKeyEnv) + } + noLeak(t, ui, err) +} + +// A refused key fails the authentication step, offers retry, edit or exit, and on exit stops with nothing saved. (SHARD-657) +func TestARefusedKeyOffersRetryEditOrExit(t *testing.T) { + url, ca := tlsFront(t, "the-accepted-key", nil) + host, path := testHost(t, map[string]string{client.CAFileEnv: ca}) + ui := &fakeUI{ + texts: map[Question]string{AskURL: url}, + secrets: map[Question]string{AskAPIKey: testKey}, + selects: map[Question]string{AskRetry: "exit"}, + } + + err := (&Setup{Host: host, UI: ui}).remote(t.Context()) + + var refused *client.APIError + if !errors.As(err, &refused) || refused.Status != http.StatusUnauthorized { + t.Fatalf("remote returned %v, want the 401", err) + } + var stopped *StoppedError + if !errors.As(err, &stopped) || stopped.Step != "Verify authentication" { + t.Errorf("remote returned %v, want it stopped at the authentication step", err) + } + want := []string{"start 0", "done 0", "start 1", "fail 1: The server did not accept the API key. / Check the key, or ask the server administrator for a new one."} + if !slices.Equal(ui.lists[0].marks, want) { + t.Errorf("marked %v, want %v", ui.lists[0].marks, want) + } + if got := names(ui, AskRetry); !slices.Equal(got, []string{"retry", "edit", "exit"}) { + t.Errorf("offered %v after a failure, want retry, edit and exit", got) + } + if got := savedConnection(t, path); got != (client.Config{}) { + t.Errorf("saved %v after a failed check", got) + } + noLeak(t, ui, err) +} + +// keysUI answers each key prompt from keys in turn, so a test can fix a wrong key through Edit. +type keysUI struct { + *fakeUI + keys []string +} + +func (u *keysUI) Secret(_ context.Context, q Question, _ string) (string, error) { + u.ask(q) + if len(u.keys) == 0 { + return "", fmt.Errorf("secret %s: %w", q, errUnscripted) + } + key := u.keys[0] + u.keys = u.keys[1:] + + return key, nil +} + +// Edit asks for the details again and checks them from the first step. (SHARD-657) +func TestEditAsksAgainAndChecksTheNewDetails(t *testing.T) { + url, ca := tlsFront(t, testKey, nil) + host, path := testHost(t, map[string]string{client.CAFileEnv: ca}) + fake := &fakeUI{ + texts: map[Question]string{AskURL: url}, + selects: map[Question]string{AskRetry: "edit"}, + confirms: map[Question]bool{AskSave: true}, + } + ui := &keysUI{fakeUI: fake, keys: []string{"a-wrong-key", testKey}} + + if err := (&Setup{Host: host, UI: ui}).remote(t.Context()); err != nil { + t.Fatalf("remote: %v", err) + } + + if want := []Question{AskURL, AskAPIKey, AskRetry, AskURL, AskAPIKey, AskSave}; !slices.Equal(fake.asked, want) { + t.Errorf("asked %v, want %v", fake.asked, want) + } + if len(fake.lists) != 2 || !slices.Equal(fake.lists[1].marks, allDone) { + t.Fatalf("drew %d checklists, want a second that passes", len(fake.lists)) + } + if got := savedConnection(t, path); got.APIKey != testKey { + t.Errorf("saved a key other than the edited one") + } +} + +// An edit asks for the key even with SHARD_API_KEY set, since that key is the one the server refused. (SHARD-657) +func TestEditAsksForTheKeyThatSHARDAPIKEYGotWrong(t *testing.T) { + url, ca := tlsFront(t, testKey, nil) + host, path := testHost(t, map[string]string{client.CAFileEnv: ca, client.APIKeyEnv: "a-stale-key"}) + ui := &fakeUI{ + texts: map[Question]string{AskURL: url}, + secrets: map[Question]string{AskAPIKey: testKey}, + selects: map[Question]string{AskRetry: "edit"}, + confirms: map[Question]bool{AskSave: true}, + } + + err := (&Setup{Host: host, UI: ui}).remote(t.Context()) + if err != nil { + t.Fatalf("remote: %v", err) + } + + if want := []Question{AskURL, AskRetry, AskURL, AskAPIKey, AskSave}; !slices.Equal(ui.asked, want) { + t.Errorf("asked %v, want %v", ui.asked, want) + } + if got := savedConnection(t, path); got.APIKey != testKey { + t.Errorf("saved a key other than the edited one") + } + containsAll(t, ui.printed, client.APIKeyEnv+" is set to another key, and Shard commands use it before the saved one.") + noLeak(t, ui, err) +} + +// Retry checks the same details again, without asking for them. (SHARD-657) +func TestRetryChecksTheSameDetailsAgain(t *testing.T) { + var requests atomic.Int32 + url, ca := tlsFront(t, testKey, func() bool { return requests.Add(1) == 1 }) + host, path := testHost(t, map[string]string{client.CAFileEnv: ca}) + ui := &fakeUI{ + texts: map[Question]string{AskURL: url}, + secrets: map[Question]string{AskAPIKey: testKey}, + selects: map[Question]string{AskRetry: "retry"}, + confirms: map[Question]bool{AskSave: true}, + } + + if err := (&Setup{Host: host, UI: ui}).remote(t.Context()); err != nil { + t.Fatalf("remote: %v", err) + } + + if want := []Question{AskURL, AskAPIKey, AskRetry, AskSave}; !slices.Equal(ui.asked, want) { + t.Errorf("asked %v, want %v", ui.asked, want) + } + if len(ui.lists) != 2 || !strings.HasPrefix(ui.lists[0].marks[3], "fail 1: the daemon is restarting") || !slices.Equal(ui.lists[1].marks, allDone) { + t.Errorf("drew %d checklists, want a 503 then a pass", len(ui.lists)) + } + if got := savedConnection(t, path); got.Remote != url { + t.Errorf("saved %v after the retry", got) + } +} + +// An untrusted certificate fails the first step, explains SHARD_CA_FILE, and offers no way around the check. (SHARD-657) +func TestAnUntrustedCertificateExplainsSHARDCAFILE(t *testing.T) { + url, _ := tlsFront(t, testKey, nil) + host, _ := testHost(t, nil) + ui := &fakeUI{ + texts: map[Question]string{AskURL: url}, + secrets: map[Question]string{AskAPIKey: testKey}, + selects: map[Question]string{AskRetry: "exit"}, + } + + err := (&Setup{Host: host, UI: ui}).remote(t.Context()) + + var untrusted *tls.CertificateVerificationError + var stopped *StoppedError + if !errors.As(err, &untrusted) || !errors.As(err, &stopped) || stopped.Step != "Reach the server" { + t.Fatalf("remote returned %v, want the certificate error, stopped at the first step", err) + } + want := []string{"start 0", "fail 0: This machine does not trust the server certificate. / For a private certificate authority, exit, set SHARD_CA_FILE to its PEM file, and run shard setup again."} + if !slices.Equal(ui.lists[0].marks, want) { + t.Errorf("marked %v, want %v", ui.lists[0].marks, want) + } + if got := names(ui, AskRetry); !slices.Equal(got, []string{"retry", "edit", "exit"}) { + t.Errorf("offered %v, want no choice beyond retry, edit and exit", got) + } + noLeak(t, ui, err) +} + +// A server nothing answers on fails the first step in the §10 shape: what failed, why, and what to check. (SHARD-657) +func TestAnUnreachableServerSaysWhyAndWhatToCheck(t *testing.T) { + server := front(testKey, nil) + server.StartTLS() + url := server.URL + server.Close() + host, _ := testHost(t, nil) + ui := &fakeUI{ + texts: map[Question]string{AskURL: url}, + secrets: map[Question]string{AskAPIKey: testKey}, + selects: map[Question]string{AskRetry: "exit"}, + } + + err := (&Setup{Host: host, UI: ui}).remote(t.Context()) + + var stopped *StoppedError + if !errors.As(err, &stopped) || stopped.Step != "Reach the server" { + t.Fatalf("remote returned %v, want it stopped at the first step", err) + } + want := []string{"start 0", "fail 0: Could not reach " + url + ": connection refused. / Check the URL, and that shard serve or the proxy in front of it runs."} + if !slices.Equal(ui.lists[0].marks, want) { + t.Errorf("marked %v, want %v", ui.lists[0].marks, want) + } +} + +// The dial failures a person meets most are worded plainly, and any other keeps the dialer's text. (SHARD-657) +func TestDialCause(t *testing.T) { + for _, tc := range []struct { + err error + want string + }{ + {&net.OpError{Op: "dial", Net: "tcp", Err: &net.DNSError{Err: "no such host", Name: "shard.invalid", IsNotFound: true}}, "no such host"}, + {&net.OpError{Op: "dial", Net: "tcp", Err: os.NewSyscallError("connect", syscall.ECONNREFUSED)}, "connection refused"}, + {&net.OpError{Op: "dial", Net: "tcp", Err: context.DeadlineExceeded}, "connection timed out"}, + {&net.OpError{Op: "dial", Net: "tcp", Err: &net.DNSError{Err: "i/o timeout", Name: "shard.example.com", IsTimeout: true}}, "connection timed out"}, + {errors.New("remote error: tls: handshake failure"), "remote error: tls: handshake failure"}, + } { + if got := dialCause(tc.err); got != tc.want { + t.Errorf("dialCause(%v) = %q, want %q", tc.err, got, tc.want) + } + } +} + +// A declined save writes nothing and says how to use the connection, with the key as a placeholder. (SHARD-657) +func TestADeclinedSaveWritesNothing(t *testing.T) { + url, ca := tlsFront(t, testKey, nil) + host, path := testHost(t, map[string]string{client.CAFileEnv: ca}) + ui := &fakeUI{ + texts: map[Question]string{AskURL: url}, + secrets: map[Question]string{AskAPIKey: testKey}, + confirms: map[Question]bool{AskSave: false}, + } + + err := (&Setup{Host: host, UI: ui}).remote(t.Context()) + if err != nil { + t.Fatalf("remote: %v", err) + } + + if _, err := os.Stat(path); !errors.Is(err, os.ErrNotExist) { + t.Errorf("stat %s returned %v after a declined save, want no file", path, err) + } + containsAll(t, ui.printed, + "The connection was verified but not saved.", + " export "+client.RemoteEnv+"="+url, + " export "+client.APIKeyEnv+"=<your API key>", + " export "+client.CAFileEnv+"="+ca, + ) + noLeak(t, ui, err) +} + +// A saved connection opens the §14 menu, and Check verifies it as saved without asking to save it again. (SHARD-657) +func TestASavedConnectionCanBeChecked(t *testing.T) { + url, ca := tlsFront(t, testKey, nil) + host, path := testHost(t, map[string]string{client.CAFileEnv: ca}) + saved := client.Config{Remote: url, APIKey: testKey} + saveConnection(t, path, saved) + ui := &fakeUI{selects: map[Question]string{AskMode: "remote", AskSaved: "check"}} + + err := (&Setup{Host: host, UI: ui}).Run(t.Context()) + if err != nil { + t.Fatalf("Run: %v", err) + } + + if ui.printed[0] != "Saved connection: "+url { + t.Errorf("printed %q first, want the saved connection", ui.printed[0]) + } + if got := names(ui, AskSaved); !slices.Equal(got, []string{"check", "replace", "remove", "exit"}) { + t.Errorf("the menu offers %v", got) + } + if want := []Question{AskMode, AskSaved}; !slices.Equal(ui.asked, want) { + t.Errorf("asked %v, want %v", ui.asked, want) + } + if len(ui.lists) != 1 || !slices.Equal(ui.lists[0].marks, allDone) { + t.Errorf("drew %d checklists, want one that passes", len(ui.lists)) + } + containsAll(t, ui.printed, "✓ Connected to "+url, " snapshot not supported") + if got := savedConnection(t, path); got != saved { + t.Errorf("a check changed the saved connection to %v", got) + } + noLeak(t, ui, err) +} + +// Replace verifies the new connection before it is saved, and a failure or a no keeps the old one. (SHARD-657) +func TestReplaceKeepsTheOldConnectionUntilTheNewOneIsSaved(t *testing.T) { + oldURL, _ := tlsFront(t, "old-key", nil) + newURL, ca := tlsFront(t, testKey, nil) + host, path := testHost(t, map[string]string{client.CAFileEnv: ca}) + old := client.Config{Remote: oldURL, APIKey: "old-key"} + saveConnection(t, path, old) + + for _, tc := range []struct { + name string + key string + save bool + want client.Config + }{ + {name: "failed check", key: "a-wrong-key", want: old}, + {name: "declined save", key: testKey, want: old}, + {name: "saved", key: testKey, save: true, want: client.Config{Remote: newURL, APIKey: testKey}}, + } { + t.Run(tc.name, func(t *testing.T) { + ui := &fakeUI{ + selects: map[Question]string{AskSaved: "replace", AskRetry: "exit"}, + texts: map[Question]string{AskURL: newURL}, + secrets: map[Question]string{AskAPIKey: tc.key}, + confirms: map[Question]bool{AskSave: tc.save}, + } + err := (&Setup{Host: host, UI: ui}).remote(t.Context()) + if tc.key == testKey && err != nil { + t.Fatalf("remote: %v", err) + } + if got := savedConnection(t, path); got != tc.want { + t.Errorf("the saved connection is %v, want %v", got, tc.want) + } + if tc.key == testKey && !tc.save { + containsAll(t, ui.printed, "The saved connection to "+oldURL+" is unchanged.") + } + noLeak(t, ui, err) + }) + } +} + +// Remove deletes the saved connection, and names SHARD_REMOTE when it still beats the local daemon. (SHARD-657) +func TestRemoveDeletesTheSavedConnection(t *testing.T) { + for _, tc := range []struct { + name string + vars map[string]string + want []string + }{ + {name: "local", want: []string{"✓ Connection removed", "", "Shard commands now use the local daemon."}}, + {name: "SHARD_REMOTE", vars: map[string]string{client.RemoteEnv: "https://other.example.com"}, want: []string{ + "✓ Connection removed", "", "SHARD_REMOTE is still set to https://other.example.com, and it overrides the local default.", "Unset it to use the local daemon.", + }}, + } { + t.Run(tc.name, func(t *testing.T) { + host, path := testHost(t, tc.vars) + saveConnection(t, path, client.Config{Remote: "https://shard.example.com", APIKey: testKey}) + ui := &fakeUI{selects: map[Question]string{AskSaved: "remove"}} + + if err := (&Setup{Host: host, UI: ui}).remote(t.Context()); err != nil { + t.Fatalf("remote: %v", err) + } + + if got := savedConnection(t, path); got != (client.Config{}) { + t.Errorf("the connection is still saved: %v", got) + } + if got := ui.printed[2:]; !slices.Equal(got, tc.want) { + t.Errorf("printed %q, want %q", got, tc.want) + } + }) + } +} + +// Exit leaves the saved connection as it is and checks nothing. (SHARD-657) +func TestExitLeavesTheSavedConnection(t *testing.T) { + host, path := testHost(t, nil) + saved := client.Config{Remote: "https://shard.example.com", APIKey: testKey} + saveConnection(t, path, saved) + ui := &fakeUI{selects: map[Question]string{AskSaved: "exit"}} + + if err := (&Setup{Host: host, UI: ui}).remote(t.Context()); err != nil { + t.Fatalf("remote: %v", err) + } + + if len(ui.lists) != 0 || savedConnection(t, path) != saved { + t.Errorf("exit checked or changed the saved connection") + } +} + +// A switch to local asks now and removes the saved connection only when finish runs, after local setup succeeds. (SHARD-657) +func TestSwitchToLocalRemovesTheConnectionOnlyAtTheEnd(t *testing.T) { + saved := client.Config{Remote: "https://shard.example.com", APIKey: testKey} + + for _, tc := range []struct { + name string + remove bool + want client.Config + says string + }{ + {name: "remove", remove: true, want: client.Config{}, says: "Shard commands now use the local daemon."}, + {name: "keep", want: saved, says: "The saved connection remains, so normal Shard commands still use the remote server."}, + } { + t.Run(tc.name, func(t *testing.T) { + host, path := testHost(t, nil) + saveConnection(t, path, saved) + ui := &fakeUI{confirms: map[Question]bool{AskSwitch: tc.remove}} + s := &Setup{Host: host, UI: ui} + + finish, err := s.switchToLocal(t.Context()) + if err != nil { + t.Fatalf("switchToLocal: %v", err) + } + containsAll(t, ui.printed, "Normal Shard commands currently use the remote server https://shard.example.com, saved in "+path+".") + if got := savedConnection(t, path); got != saved { + t.Fatalf("the connection changed before local setup finished: %v", got) + } + + if err := finish(t.Context()); err != nil { + t.Fatalf("finish: %v", err) + } + if got := savedConnection(t, path); got != tc.want { + t.Errorf("after local setup the saved connection is %v, want %v", got, tc.want) + } + containsAll(t, ui.printed, tc.says) + noLeak(t, ui, nil) + }) + } +} + +// With no saved connection a switch asks nothing, and says only that SHARD_REMOTE still beats the local daemon. (SHARD-657) +func TestSwitchToLocalWithNoSavedConnection(t *testing.T) { + for _, tc := range []struct { + name string + vars map[string]string + want []string + }{ + {name: "quiet"}, + {name: "SHARD_REMOTE", vars: map[string]string{client.RemoteEnv: "https://shard.example.com"}, want: []string{ + "SHARD_REMOTE is still set to https://shard.example.com, and it overrides the local default.", "Unset it to use the local daemon.", + }}, + } { + t.Run(tc.name, func(t *testing.T) { + host, _ := testHost(t, tc.vars) + ui := &fakeUI{} + finish, err := (&Setup{Host: host, UI: ui}).switchToLocal(t.Context()) + if err != nil { + t.Fatalf("switchToLocal: %v", err) + } + if err := finish(t.Context()); err != nil { + t.Fatalf("finish: %v", err) + } + + if len(ui.asked) != 0 || !slices.Equal(ui.printed, tc.want) { + t.Errorf("asked %v and printed %q, want nothing asked and %q", ui.asked, ui.printed, tc.want) + } + }) + } +} diff --git a/services/setup/saved.go b/services/setup/saved.go new file mode 100644 index 000000000..df0be4959 --- /dev/null +++ b/services/setup/saved.go @@ -0,0 +1,128 @@ +package setup + +import ( + "context" + "strings" + + "github.com/presmihaylov/shard/pkg/term" + "github.com/presmihaylov/shard/services/client" +) + +// The §14 choices for a saved connection, in the order the wizard shows them. +const ( + savedCheck = iota + savedReplace + savedRemove + savedExit +) + +var savedOptions = []term.Option{ + savedCheck: {Name: "check", Label: "Check this connection", Default: true}, + savedReplace: {Name: "replace", Label: "Replace this connection"}, + savedRemove: {Name: "remove", Label: "Remove this connection"}, + savedExit: {Name: "exit", Label: "Exit"}, +} + +// saved is the §14 menu for the connection saved at path. +func (s *Setup) saved(ctx context.Context, path string, saved client.Config) error { + if err := s.UI.Print("Saved connection: "+client.Redacted(saved.Remote), ""); err != nil { + return err + } + choice, err := s.UI.Select(ctx, AskSaved, "What would you like to do?", savedOptions) + if err != nil { + return err + } + + switch choice { + case savedCheck: + return s.check(ctx, path, saved) + case savedReplace: + return s.connect(ctx, path, saved) + case savedRemove: + return s.forget(path) + } + + return nil +} + +// check verifies the saved connection as it is; one edited after a failure is a replacement, saved only on the user's word. +func (s *Setup) check(ctx context.Context, path string, saved client.Config) error { + conn, caps, err := s.verify(ctx, saved) + if err != nil { + return err + } + if conn != saved { + return s.offerSave(ctx, path, saved, conn, caps) + } + + return s.UI.Print(connectedLines(conn, caps)...) +} + +// forget removes the saved connection and says what normal commands use now. +func (s *Setup) forget(path string) error { + if err := client.RemoveConfig(path); err != nil { + return err + } + lines := []string{"✓ Connection removed", ""} + if note := remoteEnvNote(s.Host.Env); note != nil { + return s.UI.Print(append(lines, note...)...) + } + + return s.UI.Print(append(lines, "Shard commands now use the local daemon.")...) +} + +// switchToLocal is §14: it asks now, and finish drops the saved connection only once local setup succeeds. +func (s *Setup) switchToLocal(ctx context.Context) (finish func(context.Context) error, err error) { + path, err := client.ConfigPath(s.Host.Env) + if err != nil { + return nil, err + } + saved, err := client.LoadConfig(path) + if err != nil { + return nil, err + } + if saved.Remote == "" { + return func(context.Context) error { return s.printNote(remoteEnvNote(s.Host.Env)) }, nil + } + + if err := s.UI.Print("Normal Shard commands currently use the remote server "+client.Redacted(saved.Remote)+", saved in "+path+".", ""); err != nil { + return nil, err + } + remove, err := s.UI.Confirm(ctx, AskSwitch, "Remove the saved connection after local setup succeeds?", true) + if err != nil { + return nil, err + } + if !remove { + return func(context.Context) error { + lines := []string{"", "The saved connection remains, so normal Shard commands still use the remote server.", "Run shard setup again to remove it."} + return s.printNote(append(lines, remoteEnvNote(s.Host.Env)...)) + }, nil + } + + return func(context.Context) error { + if err := s.UI.Print(""); err != nil { + return err + } + + return s.forget(path) + }, nil +} + +// printNote prints lines when there are any, so a run with nothing to say prints nothing. +func (s *Setup) printNote(lines []string) error { + if len(lines) == 0 { + return nil + } + + return s.UI.Print(lines...) +} + +// remoteEnvNote is the §14 reminder that SHARD_REMOTE beats the local daemon, or nothing when it is unset. +func remoteEnvNote(env func(string) string) []string { + remote := strings.TrimSpace(env(client.RemoteEnv)) + if remote == "" { + return nil + } + + return []string{client.RemoteEnv + " is still set to " + client.Redacted(remote) + ", and it overrides the local default.", "Unset it to use the local daemon."} +} diff --git a/services/setup/setup.go b/services/setup/setup.go new file mode 100644 index 000000000..fd96487e1 --- /dev/null +++ b/services/setup/setup.go @@ -0,0 +1,222 @@ +// Package setup is the one implementation of shard setup. The wizard and the flags answer the same +// questions through UI, so an automated run takes every path a person can. +package setup + +import ( + "context" + "errors" + "fmt" + "net/http" + "os" + "os/exec" + "runtime" + "strings" + + "github.com/presmihaylov/shard/pkg/term" +) + +// DataDir is the one data directory setup provisions; it never asks for another. +const DataDir = "/var/lib/shard" + +// Host is the machine setup inspects and changes. A test points Root at a temp dir and Releases at a fake server. +type Host struct { + // Root prefixes every path setup reads or writes: "/" on a real host. + Root string + OS string + Arch string + // Euid is the user setup runs as: 0 runs a privileged step as it is, any other runs it under sudo. + Euid int + // Executable is the running shard binary, and Version is its release tag. + Executable string + Version string + // Releases is the base URL of the release server. + Releases string + HTTP *http.Client + Env func(string) string + // Run runs one command as it is given; a step that needs root wraps it in sudo itself. + Run func(ctx context.Context, name string, args ...string) ([]byte, error) +} + +// releases is where setup looks up and downloads a shard release. +const releases = "https://api.github.com/repos/presmihaylov/shard/releases" + +// NewHost is this machine, as the running binary of version sees it. +func NewHost(version string) (Host, error) { + executable, err := os.Executable() + if err != nil { + return Host{}, fmt.Errorf("find the running shard binary: %w", err) + } + + return Host{ + Root: "/", + OS: runtime.GOOS, + Arch: runtime.GOARCH, + Euid: os.Geteuid(), + Executable: executable, + Version: version, + Releases: releases, + HTTP: http.DefaultClient, + Env: os.Getenv, + Run: func(ctx context.Context, name string, args ...string) ([]byte, error) { + return exec.CommandContext(ctx, name, args...).CombinedOutput() + }, + }, nil +} + +// Question names one choice, so the flags can answer it and a run without a terminal can name the flag it lacks. +type Question string + +const ( + AskMode Question = "mode" + AskProvider Question = "provider" + AskStartAtBoot Question = "start-at-boot" + AskConfirm Question = "confirm" + AskURL Question = "url" + AskHTTP Question = "http" + AskAPIKey Question = "api-key" + AskSave Question = "save" + AskSaved Question = "saved" + AskExisting Question = "existing" + AskRetry Question = "retry" + AskSwitch Question = "switch" +) + +// UI is where every answer comes from and every line goes: a person at a terminal, the flags, or a test. +type UI interface { + Select(ctx context.Context, q Question, title string, options []term.Option) (int, error) + // Confirm asks yes or no, and yes is the answer an empty reply takes. + Confirm(ctx context.Context, q Question, text string, yes bool) (bool, error) + Text(ctx context.Context, q Question, prompt string) (string, error) + // Secret reads a value that never echoes and never lands in a log. + Secret(ctx context.Context, q Question, prompt string) (string, error) + Checklist(title string, steps []string) (Checklist, error) + Print(lines ...string) error +} + +// Checklist is the live list of one job's steps; term.Checklist is the one a terminal draws. +type Checklist interface { + Start(i int) error + Done(i int) error + Attention(i int, detail ...string) error + Fail(i int, detail ...string) error +} + +// Step is one line of a checklist and the change it makes. Do is safe to run again after a stop. +type Step struct { + Title string + Do func(ctx context.Context) error +} + +// StoppedError is a step that failed: the steps before it stay in place and a second run retries. +type StoppedError struct { + Step string + Err error +} + +func (e *StoppedError) Error() string { return fmt.Sprintf("%s: %v", e.Step, e.Err) } + +func (e *StoppedError) Unwrap() error { return e.Err } + +// Problem is a step error worded for a person: apply prints its lines under the failed step. +type Problem struct { + Lines []string +} + +func (p *Problem) Error() string { return strings.Join(p.Lines, "; ") } + +// ErrDeclined is a person who said no at the confirmation, which changed nothing. +var ErrDeclined = errors.New("setup was cancelled; nothing changed") + +// Setup is one run of shard setup. +type Setup struct { + Host Host + UI UI +} + +// The first choice, in the order the wizard shows it. +const ( + modeLocal = iota + modeRemote +) + +// Run asks the first question and runs the local or the remote half. +func (s *Setup) Run(ctx context.Context) error { + mode, err := s.UI.Select(ctx, AskMode, "How do you want to use Shard?", []term.Option{ + modeLocal: {Name: "local", Label: "Run sandboxes on this machine", Default: true}, + modeRemote: {Name: "remote", Label: "Connect to a remote server"}, + }) + if err != nil { + return err + } + if mode == modeRemote { + return s.remote(ctx) + } + + return s.runLocal(ctx) +} + +// runLocal offers the existing installation's menu when there is one, else a fresh local setup. +func (s *Setup) runLocal(ctx context.Context) error { + inst, found, err := Detect(ctx, s.Host) + if err != nil { + return fmt.Errorf("look for an existing installation: %w", err) + } + if found { + return s.existing(ctx, inst) + } + + return s.switched(ctx, s.local) +} + +// switched runs a local job after the offer to drop a saved remote, and drops it only once the job succeeds. +func (s *Setup) switched(ctx context.Context, job func(context.Context) error) error { + finish, err := s.switchToLocal(ctx) + if err != nil { + return err + } + if err := job(ctx); err != nil { + return err + } + + return finish(ctx) +} + +// apply runs the steps under a live checklist, and on a failure marks the step, says what stays and stops. +func (s *Setup) apply(ctx context.Context, title string, steps []Step) error { + titles := make([]string, 0, len(steps)) + for _, step := range steps { + titles = append(titles, step.Title) + } + list, err := s.UI.Checklist(title, titles) + if err != nil { + return err + } + + for i, step := range steps { + if err := list.Start(i); err != nil { + return err + } + if err := step.Do(ctx); err != nil { + return errors.Join( + list.Fail(i, problemLines(err)...), + s.UI.Print("", "Setup stopped. Earlier completed steps remain in place.", "Run `shard setup` again to retry."), + &StoppedError{Step: step.Title, Err: err}, + ) + } + if err := list.Done(i); err != nil { + return err + } + } + + return nil +} + +// problemLines are the lines a Problem gives, or the error itself when the step returned none. +func problemLines(err error) []string { + var problem *Problem + if errors.As(err, &problem) { + return problem.Lines + } + + return []string{err.Error()} +} diff --git a/services/setup/setup_test.go b/services/setup/setup_test.go new file mode 100644 index 000000000..cf33027dc --- /dev/null +++ b/services/setup/setup_test.go @@ -0,0 +1,88 @@ +package setup + +import ( + "context" + "errors" + "fmt" + "slices" + "testing" +) + +func TestRunAsksLocalOrRemoteFirst(t *testing.T) { + ui := &fakeUI{} + err := (&Setup{UI: ui}).Run(t.Context()) + if !errors.Is(err, errUnscripted) { + t.Fatalf("Run without an answer: %v", err) + } + + if !slices.Equal(ui.asked, []Question{AskMode}) { + t.Errorf("asked %v, want only %s", ui.asked, AskMode) + } + var names []string + for _, o := range ui.options[AskMode] { + names = append(names, o.Name) + } + if !slices.Equal(names, []string{"local", "remote"}) || !ui.options[AskMode][0].Default { + t.Errorf("the first choice offers %v, want local (the default) then remote", names) + } +} + +func TestApplyMarksEveryStepDone(t *testing.T) { + ui := &fakeUI{} + ran := 0 + step := func(context.Context) error { ran++; return nil } + + if err := (&Setup{UI: ui}).apply(t.Context(), "Setting up", []Step{{"one", step}, {"two", step}}); err != nil { + t.Fatal(err) + } + + if ran != 2 || !slices.Equal(ui.lists[0].steps, []string{"one", "two"}) { + t.Errorf("ran %d steps of %v", ran, ui.lists[0].steps) + } + if want := []string{"start 0", "done 0", "start 1", "done 1"}; !slices.Equal(ui.lists[0].marks, want) { + t.Errorf("marked %v, want %v", ui.lists[0].marks, want) + } + if len(ui.printed) > 0 { + t.Errorf("a clean run printed %q", ui.printed) + } +} + +func TestApplyStopsAtTheFailedStepAndSaysWhatStays(t *testing.T) { + ui := &fakeUI{} + broke := errors.New("it broke") + ran := []string{} + step := func(name string, err error) Step { + return Step{name, func(context.Context) error { ran = append(ran, name); return err }} + } + + err := (&Setup{UI: ui}).apply(t.Context(), "Setting up", []Step{step("one", nil), step("two", broke), step("three", nil)}) + + var stopped *StoppedError + if !errors.As(err, &stopped) || stopped.Step != "two" || !errors.Is(err, broke) { + t.Fatalf("apply returned %v, want the stop at two", err) + } + if !slices.Equal(ran, []string{"one", "two"}) { + t.Errorf("ran %v, want it to stop after two", ran) + } + if want := []string{"start 0", "done 0", "start 1", "fail 1: it broke"}; !slices.Equal(ui.lists[0].marks, want) { + t.Errorf("marked %v, want %v", ui.lists[0].marks, want) + } + if want := []string{"", "Setup stopped. Earlier completed steps remain in place.", "Run `shard setup` again to retry."}; !slices.Equal(ui.printed, want) { + t.Errorf("printed %q, want %q", ui.printed, want) + } +} + +func TestApplyPrintsTheLinesOfAProblem(t *testing.T) { + ui := &fakeUI{} + problem := &Problem{Lines: []string{"Could not download runsc", "Check the network and run shard setup again."}} + fail := func(context.Context) error { return fmt.Errorf("install gVisor: %w", problem) } + + err := (&Setup{UI: ui}).apply(t.Context(), "Setting up", []Step{{"Install gVisor", fail}}) + + if !errors.Is(err, problem) { + t.Fatalf("apply returned %v, want the problem", err) + } + if want := []string{"start 0", "fail 0: Could not download runsc / Check the network and run shard setup again."}; !slices.Equal(ui.lists[0].marks, want) { + t.Errorf("marked %v, want %v", ui.lists[0].marks, want) + } +} diff --git a/services/setup/tools.go b/services/setup/tools.go new file mode 100644 index 000000000..23cd41429 --- /dev/null +++ b/services/setup/tools.go @@ -0,0 +1,156 @@ +package setup + +import ( + "os" + "path/filepath" + "slices" + "strings" +) + +// tool is one program a provider's daemon runs, which setup finds where the root daemon looks or installs. +type tool struct { + // Name is the command the daemon runs. + Name string + // Package is the apt package that gives it, and "" when a pinned download does or setup never installs it. + Package string + // Download is the pinned release that gives it. + Download *download +} + +// download is one pinned upstream release; setup checks its sha256 before it installs any file of it. +type download struct { + // Title names it in a message, as the tool it gives. + Title string + URL string + SHA256 string + // Files maps an entry of the archive to the path it installs at; a .deb has none, apt installs it whole. + Files map[string]string +} + +func (d *download) deb() bool { return strings.HasSuffix(d.URL, ".deb") } + +// The pins: bump each one with its sha256 from the release page, never by name alone. +var ( + gvisorRelease = &download{ + Title: "runsc", + URL: "https://storage.googleapis.com/gvisor/releases/release/20260928.0/x86_64/gvisor.tar.zstd", + SHA256: "1f2a732d12072eda6b8a64eccf9829fe9f4a1893d3292a7e8005a82089811d8d", + // runsc execs its sidecars from gvisor-bin beside itself, so they move together; the containerd shim is for containerd only. + Files: map[string]string{ + "runsc": "/usr/local/bin/runsc", + "gvisor-bin/checkpointgofer": "/usr/local/bin/gvisor-bin/checkpointgofer", + "gvisor-bin/gvisor-sentry-prewarmer": "/usr/local/bin/gvisor-bin/gvisor-sentry-prewarmer", + "gvisor-bin/gvisor_sentry": "/usr/local/bin/gvisor-bin/gvisor_sentry", + "gvisor-bin/runsc-fd-parking": "/usr/local/bin/gvisor-bin/runsc-fd-parking", + "gvisor-bin/runsc-metric-server": "/usr/local/bin/gvisor-bin/runsc-metric-server", + }, + } + firecrackerRelease = &download{ + Title: "firecracker", + URL: "https://github.com/firecracker-microvm/firecracker/releases/download/v1.17.0/firecracker-v1.17.0-x86_64.tgz", + SHA256: "06094a1108ae9e82aa4c23a775aa92758f53f1175d422270d9d6162cb9ade558", + Files: map[string]string{ + "release-v1.17.0-x86_64/firecracker-v1.17.0-x86_64": "/usr/local/bin/firecracker", + "release-v1.17.0-x86_64/jailer-v1.17.0-x86_64": "/usr/local/bin/jailer", + }, + } + sysboxRelease = &download{ + Title: "sysbox", + URL: "https://github.com/nestybox/sysbox/releases/download/v0.7.1/sysbox-ce_0.7.1.linux_amd64.deb", + SHA256: "9d6d5484f980d0a17f86c492c1262015c2afb66280bdb97215b79fde6a0261c5", + } +) + +// linuxTools are what every Linux provider runs: the sandbox network, its policy and the sandbox disks. +var linuxTools = []tool{ + {Name: "ip", Package: "iproute2"}, + {Name: "nft", Package: "nftables"}, + {Name: "mkfs.ext4", Package: "e2fsprogs"}, +} + +// providerTools are what the daemon of each provider runs beyond linuxTools. +var providerTools = map[string][]tool{ + Firecracker: { + {Name: "firecracker", Download: firecrackerRelease}, + {Name: "jailer", Download: firecrackerRelease}, + {Name: "mkfs.erofs", Package: "erofs-utils"}, + {Name: "mkfs.xfs", Package: "xfsprogs"}, + }, + GVisor: {{Name: "runsc", Download: gvisorRelease}}, + Sysbox: {{Name: "sysbox-runc", Download: sysboxRelease}}, + Runc: {{Name: "runc", Package: "runc"}}, + // codesign signs the VM shim on first use, and it ships with macOS. + VZ: {{Name: "codesign"}}, +} + +// apparmorTool loads the profile runc confines a sandbox with, on a kernel that enforces AppArmor. +var apparmorTool = tool{Name: "apparmor_parser", Package: "apparmor"} + +// toolsFor are the tools the daemon of provider runs on this host. +func toolsFor(h Host, provider string) []tool { + var tools []tool + if provider != VZ { + tools = append(tools, linuxTools...) + } + tools = append(tools, providerTools[provider]...) + if provider == Runc && apparmorEnabled(h) { + tools = append(tools, apparmorTool) + } + + return tools +} + +func apparmorEnabled(h Host) bool { + enabled, err := os.ReadFile(filepath.Join(h.Root, "sys", "module", "apparmor", "parameters", "enabled")) + + return err == nil && strings.HasPrefix(string(enabled), "Y") +} + +// missing is what setup has to install for a provider: the absent tools, the packages and the downloads that give them. +type missing struct { + tools []tool + packages []string + downloads []*download + // manual are absent tools setup cannot install. + manual []string +} + +func missingTools(h Host, provider string) missing { + var m missing + for _, t := range toolsFor(h, provider) { + if _, ok := lookPath(h, t.Name); ok { + continue + } + m.tools = append(m.tools, t) + switch { + case t.Download != nil: + if !slices.Contains(m.downloads, t.Download) { + m.downloads = append(m.downloads, t.Download) + } + case t.Package != "": + m.packages = append(m.packages, t.Package) + default: + m.manual = append(m.manual, t.Name) + } + } + + return m +} + +func (m missing) names() []string { + names := make([]string, 0, len(m.tools)) + for _, t := range m.tools { + names = append(names, t.Name) + } + + return names +} + +// apt is whether installing m runs apt-get: a package, or a .deb it resolves the dependencies of. +func (m missing) apt() bool { + if len(m.packages) > 0 { + return true + } + + return slices.ContainsFunc(m.downloads, (*download).deb) +} diff --git a/services/setup/ui_test.go b/services/setup/ui_test.go new file mode 100644 index 000000000..c5fd7dbf0 --- /dev/null +++ b/services/setup/ui_test.go @@ -0,0 +1,117 @@ +package setup + +import ( + "context" + "errors" + "fmt" + "strings" + + "github.com/presmihaylov/shard/pkg/term" +) + +// fakeUI answers each question from a script and records what it was shown, so a test drives one path of setup. +type fakeUI struct { + // selects names the option each question picks, by its Name. + selects map[Question]string + confirms map[Question]bool + texts map[Question]string + secrets map[Question]string + + asked []Question + options map[Question][]term.Option + printed []string + lists []*fakeChecklist +} + +var errUnscripted = errors.New("the test gave no answer") + +func (f *fakeUI) ask(q Question) { f.asked = append(f.asked, q) } + +func (f *fakeUI) Select(_ context.Context, q Question, _ string, options []term.Option) (int, error) { + f.ask(q) + if f.options == nil { + f.options = map[Question][]term.Option{} + } + f.options[q] = options + name, ok := f.selects[q] + if !ok { + return 0, fmt.Errorf("select %s: %w", q, errUnscripted) + } + for i, o := range options { + if o.Name == name && len(o.Unavailable) == 0 { + return i, nil + } + } + + return 0, fmt.Errorf("select %s: no option %q can be chosen", q, name) +} + +func (f *fakeUI) Confirm(_ context.Context, q Question, _ string, _ bool) (bool, error) { + f.ask(q) + answer, ok := f.confirms[q] + if !ok { + return false, fmt.Errorf("confirm %s: %w", q, errUnscripted) + } + + return answer, nil +} + +func (f *fakeUI) Text(_ context.Context, q Question, _ string) (string, error) { + f.ask(q) + answer, ok := f.texts[q] + if !ok { + return "", fmt.Errorf("text %s: %w", q, errUnscripted) + } + + return answer, nil +} + +func (f *fakeUI) Secret(_ context.Context, q Question, _ string) (string, error) { + f.ask(q) + answer, ok := f.secrets[q] + if !ok { + return "", fmt.Errorf("secret %s: %w", q, errUnscripted) + } + + return answer, nil +} + +func (f *fakeUI) Checklist(title string, steps []string) (Checklist, error) { + list := &fakeChecklist{title: title, steps: steps} + f.lists = append(f.lists, list) + + return list, nil +} + +func (f *fakeUI) Print(lines ...string) error { + f.printed = append(f.printed, lines...) + + return nil +} + +// fakeChecklist records each mark as one line, such as "done 0" or "fail 1: it broke". +type fakeChecklist struct { + title string + steps []string + marks []string +} + +func (c *fakeChecklist) mark(verb string, i int, detail []string) error { + line := fmt.Sprintf("%s %d", verb, i) + if len(detail) > 0 { + line += ": " + strings.Join(detail, " / ") + } + c.marks = append(c.marks, line) + + return nil +} + +func (c *fakeChecklist) Start(i int) error { return c.mark("start", i, nil) } + +func (c *fakeChecklist) Done(i int) error { return c.mark("done", i, nil) } + +func (c *fakeChecklist) Attention(i int, detail ...string) error { + return c.mark("attention", i, detail) +} + +func (c *fakeChecklist) Fail(i int, detail ...string) error { return c.mark("fail", i, detail) }