From 53fc4622629681a7b14de1c8cb26edd15bc5ecbd Mon Sep 17 00:00:00 2001 From: Raj Nakarja Date: Thu, 20 Aug 2026 10:06:06 +0200 Subject: [PATCH 1/3] Thread a session instead of reaching for package state Every command now takes a Session carrying the base url, version, http client, input, output and browser opener, so nothing reads or writes package-level state. TakeServerFlag returns the base rather than assigning it, and the eight mutable package vars are gone. Command matching, help rendering and the help and version answers move from main.go into internal/commands/dispatch.go, leaving main.go holding the version, the table and the wiring. Two shared helpers replace 21 copies of the server error relay and five copies of the --json argument split. The relay bounds its read, which member list did not. Polling now treats five seconds as the default for an absent interval rather than a floor over a provided one, and slow_down always adds five seconds, per RFC 8628 sections 3.2 and 3.5. --- internal/commands/account_balance.go | 28 +-- internal/commands/account_balance_test.go | 8 +- internal/commands/account_delete.go | 19 +- internal/commands/account_delete_test.go | 10 +- internal/commands/account_topup.go | 19 +- internal/commands/account_topup_test.go | 12 +- internal/commands/balances.go | 12 +- internal/commands/browser.go | 4 +- internal/commands/browser_test.go | 19 -- internal/commands/client.go | 77 ++++--- internal/commands/client_test.go | 80 +++---- internal/commands/device_claim.go | 18 +- internal/commands/device_claim_test.go | 22 +- internal/commands/device_list.go | 29 +-- internal/commands/device_list_test.go | 16 +- internal/commands/device_release.go | 24 +- internal/commands/device_release_test.go | 16 +- internal/commands/device_rename.go | 13 +- internal/commands/device_rename_test.go | 14 +- internal/commands/devices.go | 12 +- internal/commands/dispatch.go | 163 ++++++++++++++ internal/commands/dispatch_test.go | 94 ++++++++ internal/commands/fleet_create.go | 14 +- internal/commands/fleet_create_test.go | 14 +- internal/commands/fleet_delete.go | 26 +-- internal/commands/fleet_delete_test.go | 59 ++--- internal/commands/fleet_list.go | 23 +- internal/commands/fleet_list_test.go | 8 +- internal/commands/fleet_rename.go | 13 +- internal/commands/fleet_rename_test.go | 10 +- internal/commands/fleet_transfer.go | 14 +- internal/commands/fleet_transfer_test.go | 10 +- internal/commands/fleets.go | 12 +- internal/commands/fleets_test.go | 6 +- internal/commands/key_create.go | 14 +- internal/commands/key_create_test.go | 8 +- internal/commands/key_list.go | 28 +-- internal/commands/key_list_test.go | 8 +- internal/commands/key_revoke.go | 22 +- internal/commands/key_revoke_test.go | 9 +- internal/commands/keys.go | 12 +- internal/commands/login.go | 67 +++--- internal/commands/login_test.go | 86 +++----- internal/commands/logout.go | 21 +- internal/commands/logout_test.go | 17 +- internal/commands/member_add.go | 14 +- internal/commands/member_add_test.go | 10 +- internal/commands/member_list.go | 40 +--- internal/commands/member_list_test.go | 8 +- internal/commands/member_remove.go | 22 +- internal/commands/member_remove_test.go | 12 +- internal/commands/session.go | 35 +++ main.go | 255 +++++----------------- main_test.go | 106 ++------- 54 files changed, 783 insertions(+), 929 deletions(-) delete mode 100644 internal/commands/browser_test.go create mode 100644 internal/commands/dispatch.go create mode 100644 internal/commands/dispatch_test.go create mode 100644 internal/commands/session.go diff --git a/internal/commands/account_balance.go b/internal/commands/account_balance.go index 643d184..4d99031 100644 --- a/internal/commands/account_balance.go +++ b/internal/commands/account_balance.go @@ -4,24 +4,12 @@ import ( "encoding/json" "errors" "fmt" - "os" "strconv" ) -func AccountBalance(arguments []string) error { +func AccountBalance(session Session, arguments []string) error { - jsonOutput := false - - positionals := []string{} - - for _, argument := range arguments { - if argument == "--json" { - jsonOutput = true - continue - } - - positionals = append(positionals, argument) - } + positionals, jsonOutput := takeJsonFlag(arguments) if len(positionals) > 1 { return errors.New("account balance takes at most one fleet id") @@ -39,7 +27,7 @@ func AccountBalance(arguments []string) error { chosenFleetId = parsed } - fleets, err := fetchFleets() + fleets, err := fetchFleets(session) if err != nil { return err @@ -57,7 +45,7 @@ func AccountBalance(arguments []string) error { } } - fetched, err := fetchBalances() + fetched, err := fetchBalances(session) if err != nil { return err @@ -72,11 +60,11 @@ func AccountBalance(arguments []string) error { } if jsonOutput { - return json.NewEncoder(os.Stdout).Encode(balances) + return json.NewEncoder(session.Out).Encode(balances) } if len(balances) == 0 { - fmt.Println("No fleets yet. Create one with fleet create.") + fmt.Fprintln(session.Out, "No fleets yet. Create one with fleet create.") return nil } @@ -88,10 +76,10 @@ func AccountBalance(arguments []string) error { nameWidth = max(nameWidth, len(fleetNames[balance.Fleet])) } - fmt.Printf("%-*s %-*s %s\n", idWidth, "ID", nameWidth, "NAME", "BALANCE") + fmt.Fprintf(session.Out, "%-*s %-*s %s\n", idWidth, "ID", nameWidth, "NAME", "BALANCE") for _, balance := range balances { - fmt.Printf("%-*d %-*s %s\n", idWidth, balance.Fleet, nameWidth, fleetNames[balance.Fleet], formatBalance(balance)) + fmt.Fprintf(session.Out, "%-*d %-*s %s\n", idWidth, balance.Fleet, nameWidth, fleetNames[balance.Fleet], formatBalance(balance)) } return nil diff --git a/internal/commands/account_balance_test.go b/internal/commands/account_balance_test.go index 8686d0b..e176130 100644 --- a/internal/commands/account_balance_test.go +++ b/internal/commands/account_balance_test.go @@ -92,11 +92,11 @@ func TestAccountBalance(t *testing.T) { fmt.Fprint(w, test.balances) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - printed, err := captureStdout(t, func() error { - return AccountBalance(test.arguments) - }) + err := AccountBalance(session, test.arguments) + + printed := out.String() if test.wantError != "" { if err == nil || !strings.Contains(err.Error(), test.wantError) { diff --git a/internal/commands/account_delete.go b/internal/commands/account_delete.go index fba83bf..9b67ab1 100644 --- a/internal/commands/account_delete.go +++ b/internal/commands/account_delete.go @@ -4,37 +4,36 @@ import ( "bufio" "errors" "fmt" - "io" "io/fs" "net/http" "os" "strings" ) -func AccountDelete(arguments []string) error { +func AccountDelete(session Session, arguments []string) error { if len(arguments) != 0 { return errors.New("account delete takes no arguments") } - fmt.Print("Delete your account, its logins, and your access to every fleet? This cannot be undone. [y/N] ") + fmt.Fprint(session.Out, "Delete your account, its logins, and your access to every fleet? This cannot be undone. [y/N] ") - answer, _ := bufio.NewReader(os.Stdin).ReadString('\n') + answer, _ := bufio.NewReader(session.In).ReadString('\n') answer = strings.ToLower(strings.TrimSpace(answer)) if answer != "y" && answer != "yes" { - fmt.Println("Nothing deleted.") + fmt.Fprintln(session.Out, "Nothing deleted.") return nil } - request, err := authenticatedRequest(http.MethodDelete, "/account", nil) + request, err := authenticatedRequest(session, http.MethodDelete, "/account", nil) if err != nil { return err } - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return fmt.Errorf("the server could not be reached: %w", err) @@ -43,9 +42,7 @@ func AccountDelete(arguments []string) error { defer response.Body.Close() if response.StatusCode != http.StatusNoContent { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return serverError(response) } // The stored login died with the account, so it goes whether or not the @@ -62,7 +59,7 @@ func AccountDelete(arguments []string) error { return err } - fmt.Println("Account deleted.") + fmt.Fprintln(session.Out, "Account deleted.") return nil } diff --git a/internal/commands/account_delete_test.go b/internal/commands/account_delete_test.go index 734dad5..84f7748 100644 --- a/internal/commands/account_delete_test.go +++ b/internal/commands/account_delete_test.go @@ -76,7 +76,7 @@ func TestAccountDelete(t *testing.T) { w.WriteHeader(http.StatusNoContent) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) path, err := keyPath() @@ -84,11 +84,11 @@ func TestAccountDelete(t *testing.T) { t.Fatal(err) } - answerOnStdin(t, test.answer) + session.In = strings.NewReader(test.answer) - printed, err := captureStdout(t, func() error { - return AccountDelete(test.arguments) - }) + err = AccountDelete(session, test.arguments) + + printed := out.String() if test.wantError != "" { if err == nil || !strings.Contains(err.Error(), test.wantError) { diff --git a/internal/commands/account_topup.go b/internal/commands/account_topup.go index ad79c97..c50afc6 100644 --- a/internal/commands/account_topup.go +++ b/internal/commands/account_topup.go @@ -5,14 +5,11 @@ import ( "encoding/json" "errors" "fmt" - "io" "net/http" - "os" "strconv" - "strings" ) -func AccountTopup(arguments []string) error { +func AccountTopup(session Session, arguments []string) error { if len(arguments) != 1 { return errors.New("account topup takes a fleet id") @@ -24,14 +21,14 @@ func AccountTopup(arguments []string) error { return errors.New("the fleet id is the number shown by fleet list") } - request, err := authenticatedRequest(http.MethodPost, + request, err := authenticatedRequest(session, http.MethodPost, "/fleets/"+strconv.FormatInt(fleetId, 10)+"/topup", nil) if err != nil { return err } - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return fmt.Errorf("the server could not be reached: %w", err) @@ -40,9 +37,7 @@ func AccountTopup(arguments []string) error { defer response.Body.Close() if response.StatusCode != http.StatusOK { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return serverError(response) } opened := struct { @@ -55,15 +50,15 @@ func AccountTopup(arguments []string) error { return errors.New("the payment page could not be opened, try again") } - fmt.Printf("Open this link to choose an amount and pay:\n\n %s\n\nThe credit appears on the balance once the payment completes.\nPress enter to open the browser.\n", opened.Url) + fmt.Fprintf(session.Out, "Open this link to choose an amount and pay:\n\n %s\n\nThe credit appears on the balance once the payment completes.\nPress enter to open the browser.\n", opened.Url) - _, err = bufio.NewReader(os.Stdin).ReadString('\n') + _, err = bufio.NewReader(session.In).ReadString('\n') if err != nil { return nil } - openBrowser(opened.Url) + session.OpenBrowser(opened.Url) return nil } diff --git a/internal/commands/account_topup_test.go b/internal/commands/account_topup_test.go index cb028a1..a18d151 100644 --- a/internal/commands/account_topup_test.go +++ b/internal/commands/account_topup_test.go @@ -70,15 +70,15 @@ func TestAccountTopup(t *testing.T) { fmt.Fprint(w, `{"url":"https://checkout.stripe.com/c/pay/cs_test_1"}`) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) + session.In = strings.NewReader(test.stdin) - answerOnStdin(t, test.stdin) + browserOpens := make(chan string, 1) + session.OpenBrowser = func(url string) { browserOpens <- url } - browserOpens := captureBrowserOpens(t) + err := AccountTopup(session, test.arguments) - printed, err := captureStdout(t, func() error { - return AccountTopup(test.arguments) - }) + printed := out.String() if test.wantError != "" { if err == nil || !strings.Contains(err.Error(), test.wantError) { diff --git a/internal/commands/balances.go b/internal/commands/balances.go index fb24614..b3b3210 100644 --- a/internal/commands/balances.go +++ b/internal/commands/balances.go @@ -3,10 +3,8 @@ package commands import ( "encoding/json" "fmt" - "io" "net/http" "strconv" - "strings" ) type balanceEntry struct { @@ -15,14 +13,14 @@ type balanceEntry struct { Currency string `json:"currency"` } -func fetchBalances() ([]balanceEntry, error) { - request, err := authenticatedRequest(http.MethodGet, "/balance", nil) +func fetchBalances(session Session) ([]balanceEntry, error) { + request, err := authenticatedRequest(session, http.MethodGet, "/balance", nil) if err != nil { return nil, err } - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return nil, fmt.Errorf("the server could not be reached: %w", err) @@ -31,9 +29,7 @@ func fetchBalances() ([]balanceEntry, error) { defer response.Body.Close() if response.StatusCode != http.StatusOK { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return nil, fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return nil, serverError(response) } balances := []balanceEntry{} diff --git a/internal/commands/browser.go b/internal/commands/browser.go index d5bfab8..9a4887a 100644 --- a/internal/commands/browser.go +++ b/internal/commands/browser.go @@ -5,9 +5,7 @@ import ( "runtime" ) -// A variable so tests can swap in a recorder instead of reaching a real -// browser. Opening is best effort: the link is already on screen. -var openBrowser = func(url string) { +func openBrowser(url string) { switch runtime.GOOS { case "darwin": exec.Command("open", url).Start() diff --git a/internal/commands/browser_test.go b/internal/commands/browser_test.go deleted file mode 100644 index de7b52a..0000000 --- a/internal/commands/browser_test.go +++ /dev/null @@ -1,19 +0,0 @@ -package commands - -import ( - "testing" -) - -func captureBrowserOpens(t *testing.T) chan string { - t.Helper() - - opens := make(chan string, 8) - - previousOpenBrowser := openBrowser - - openBrowser = func(url string) { opens <- url } - - t.Cleanup(func() { openBrowser = previousOpenBrowser }) - - return opens -} diff --git a/internal/commands/client.go b/internal/commands/client.go index 1a36e3c..59b872b 100644 --- a/internal/commands/client.go +++ b/internal/commands/client.go @@ -10,40 +10,31 @@ import ( "path/filepath" "runtime" "strings" - "time" ) -const defaultApiBase = "https://supernext.siliconwitchery.com" - -var CliVersion = "unknown" - -var chosenApiBase = "" - -var apiClient = &http.Client{Timeout: 30 * time.Second} - // Deliberately absent from the help: development needs to point a command at // another server, users never do. -func TakeServerFlag(arguments []string) ([]string, error) { +func TakeServerFlag(arguments []string) ([]string, string, error) { remaining := []string{} - server := "" + base := defaultApiBase for index := 0; index < len(arguments); index++ { switch { case arguments[index] == "--server": if index+1 == len(arguments) || arguments[index+1] == "" { - return nil, errors.New("--server needs a url") + return nil, "", errors.New("--server needs a url") } index++ - server = arguments[index] + base = arguments[index] case strings.HasPrefix(arguments[index], "--server="): - server = strings.TrimPrefix(arguments[index], "--server=") + base = strings.TrimPrefix(arguments[index], "--server=") - if server == "" { - return nil, errors.New("--server needs a url") + if base == "" { + return nil, "", errors.New("--server needs a url") } default: @@ -51,19 +42,17 @@ func TakeServerFlag(arguments []string) ([]string, error) { } } - chosenApiBase = strings.TrimSuffix(server, "/") - - return remaining, nil + return remaining, strings.TrimSuffix(base, "/"), nil } -func CheckServer() error { - request, err := apiRequest(http.MethodGet, "/", nil) +func CheckServer(session Session) error { + request, err := apiRequest(session, http.MethodGet, "/", nil) if err != nil { return err } - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return fmt.Errorf("the server at %s cannot be reached", strings.TrimSuffix(request.URL.String(), "/")) @@ -74,25 +63,19 @@ func CheckServer() error { return nil } -func apiRequest(method string, path string, body io.Reader) (*http.Request, error) { - base := chosenApiBase - - if base == "" { - base = defaultApiBase - } - - request, err := http.NewRequest(method, strings.TrimSuffix(base, "/")+path, body) +func apiRequest(session Session, method string, path string, body io.Reader) (*http.Request, error) { + request, err := http.NewRequest(method, strings.TrimSuffix(session.Base, "/")+path, body) if err != nil { return nil, err } - request.Header.Set("User-Agent", "superstack/"+CliVersion) + request.Header.Set("User-Agent", "superstack/"+session.Version) return request, nil } -func authenticatedRequest(method string, path string, body io.Reader) (*http.Request, error) { +func authenticatedRequest(session Session, method string, path string, body io.Reader) (*http.Request, error) { storedKeyPath, err := keyPath() if err != nil { @@ -115,7 +98,7 @@ func authenticatedRequest(method string, path string, body io.Reader) (*http.Req return nil, errors.New("you are not logged in, run login first") } - request, err := apiRequest(method, path, body) + request, err := apiRequest(session, method, path, body) if err != nil { return nil, err @@ -126,6 +109,34 @@ func authenticatedRequest(method string, path string, body io.Reader) (*http.Req return request, nil } +func serverError(response *http.Response) error { + message, err := io.ReadAll(io.LimitReader(response.Body, 4096)) + + detail := strings.TrimSpace(string(message)) + + if err != nil || detail == "" { + detail = response.Status + } + + return fmt.Errorf("the server said: %s", detail) +} + +func takeJsonFlag(arguments []string) ([]string, bool) { + positionals := []string{} + jsonOutput := false + + for _, argument := range arguments { + if argument == "--json" { + jsonOutput = true + continue + } + + positionals = append(positionals, argument) + } + + return positionals, jsonOutput +} + func keyPath() (string, error) { // The key is state, not configuration: linux dotfile repos routinely // publish all of ~/.config, so the key must never live there. The mac diff --git a/internal/commands/client_test.go b/internal/commands/client_test.go index 68de4af..47494d2 100644 --- a/internal/commands/client_test.go +++ b/internal/commands/client_test.go @@ -1,7 +1,7 @@ package commands import ( - "io" + "bytes" "net/http" "net/http/httptest" "os" @@ -11,34 +11,6 @@ import ( "testing" ) -func captureStdout(t *testing.T, run func() error) (string, error) { - t.Helper() - - readEnd, writeEnd, err := os.Pipe() - - if err != nil { - t.Fatal(err) - } - - stdout := os.Stdout - - os.Stdout = writeEnd - - runError := run() - - os.Stdout = stdout - - writeEnd.Close() - - printed, err := io.ReadAll(readEnd) - - if err != nil { - t.Fatal(err) - } - - return string(printed), runError -} - func isolateKeyStorage(t *testing.T) string { t.Helper() @@ -51,7 +23,7 @@ func isolateKeyStorage(t *testing.T) string { return temporary } -func loggedInTestServer(t *testing.T, handler http.Handler) { +func loggedInSession(t *testing.T, handler http.Handler) (Session, *bytes.Buffer) { t.Helper() isolateKeyStorage(t) @@ -89,9 +61,11 @@ func loggedInTestServer(t *testing.T, handler http.Handler) { t.Cleanup(server.Close) - chosenApiBase = server.URL + out := &bytes.Buffer{} + session := NewSession(server.URL, "test", strings.NewReader(""), out) + session.OpenBrowser = func(url string) {} - t.Cleanup(func() { chosenApiBase = "" }) + return session, out } func TestKeyPathStaysOutOfPublishedDotfiles(t *testing.T) { @@ -200,11 +174,7 @@ func TestTakeServerFlag(t *testing.T) { for _, test := range tests { t.Run(test.name, func(t *testing.T) { - chosenApiBase = "" - - t.Cleanup(func() { chosenApiBase = "" }) - - remaining, err := TakeServerFlag(test.arguments) + remaining, base, err := TakeServerFlag(test.arguments) if test.wantError != "" { if err == nil || !strings.Contains(err.Error(), test.wantError) { @@ -222,20 +192,20 @@ func TestTakeServerFlag(t *testing.T) { t.Errorf("remaining = %q, want %q", strings.Join(remaining, " "), test.wantRemaining) } - if chosenApiBase != test.wantBase { - t.Errorf("chosenApiBase = %q, want %q", chosenApiBase, test.wantBase) + wantBase := test.wantBase + + if wantBase == "" { + wantBase = defaultApiBase + } + + if base != wantBase { + t.Errorf("base = %q, want %q", base, wantBase) } }) } } func TestApiRequestBase(t *testing.T) { - previousVersion := CliVersion - - CliVersion = "1.2.3" - - t.Cleanup(func() { CliVersion = previousVersion }) - tests := []struct { name string chosenBase string @@ -254,11 +224,15 @@ func TestApiRequestBase(t *testing.T) { for _, test := range tests { t.Run(test.name, func(t *testing.T) { - chosenApiBase = test.chosenBase + base := test.chosenBase - t.Cleanup(func() { chosenApiBase = "" }) + if base == "" { + base = defaultApiBase + } + + session := NewSession(base, "1.2.3", strings.NewReader(""), &bytes.Buffer{}) - request, err := apiRequest(http.MethodGet, "/login", nil) + request, err := apiRequest(session, http.MethodGet, "/login", nil) if err != nil { t.Fatal(err) @@ -282,11 +256,9 @@ func TestCheckServer(t *testing.T) { defer reachable.Close() - chosenApiBase = reachable.URL - - t.Cleanup(func() { chosenApiBase = "" }) + session := NewSession(reachable.URL, "test", strings.NewReader(""), &bytes.Buffer{}) - err := CheckServer() + err := CheckServer(session) if err != nil { t.Fatalf("a reachable server reported: %v", err) @@ -296,9 +268,9 @@ func TestCheckServer(t *testing.T) { unreachable.Close() - chosenApiBase = unreachable.URL + session.Base = unreachable.URL - err = CheckServer() + err = CheckServer(session) if err == nil || !strings.Contains(err.Error(), "cannot be reached") { t.Fatalf("error = %v, want the consistent cannot-be-reached message", err) diff --git a/internal/commands/device_claim.go b/internal/commands/device_claim.go index 27f82bc..a0e9d11 100644 --- a/internal/commands/device_claim.go +++ b/internal/commands/device_claim.go @@ -5,16 +5,14 @@ import ( "encoding/json" "errors" "fmt" - "io" "net/http" "strconv" - "strings" "time" ) -var claimClient = &http.Client{Timeout: 90 * time.Second} +func DeviceClaim(session Session, arguments []string) error { + claimClient := &http.Client{Timeout: 90 * time.Second} -func DeviceClaim(arguments []string) error { if len(arguments) != 2 && len(arguments) != 3 { return errors.New("device claim takes an IMEI, a fleet id, and an optional name") } @@ -31,7 +29,7 @@ func DeviceClaim(arguments []string) error { return errors.New("the fleet id is the number shown by fleet list") } - fleets, err := fetchFleets() + fleets, err := fetchFleets(session) if err != nil { return err @@ -49,7 +47,7 @@ func DeviceClaim(arguments []string) error { return errors.New("the fleet id is the number shown by fleet list") } - fmt.Println("Press the button on the device to finish claiming it.") + fmt.Fprintln(session.Out, "Press the button on the device to finish claiming it.") payload := map[string]string{"imei": imei} @@ -63,7 +61,7 @@ func DeviceClaim(arguments []string) error { return err } - request, err := authenticatedRequest(http.MethodPost, + request, err := authenticatedRequest(session, http.MethodPost, "/fleets/"+strconv.FormatInt(fleetId, 10)+"/devices", bytes.NewReader(body)) if err != nil { @@ -81,12 +79,10 @@ func DeviceClaim(arguments []string) error { defer response.Body.Close() if response.StatusCode != http.StatusNoContent { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return serverError(response) } - fmt.Printf("Claimed the device into %q.\n", fleetName) + fmt.Fprintf(session.Out, "Claimed the device into %q.\n", fleetName) return nil } diff --git a/internal/commands/device_claim_test.go b/internal/commands/device_claim_test.go index 52dc69c..6938199 100644 --- a/internal/commands/device_claim_test.go +++ b/internal/commands/device_claim_test.go @@ -62,11 +62,11 @@ func TestDeviceClaim(t *testing.T) { w.WriteHeader(test.statusCode) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - printed, err := captureStdout(t, func() error { - return DeviceClaim([]string{"354820091234567", "3", "roof sensor"}) - }) + err := DeviceClaim(session, []string{"354820091234567", "3", "roof sensor"}) + + printed := out.String() if test.wantError == "" && err != nil { t.Fatal(err) @@ -101,9 +101,9 @@ func TestDeviceClaimOmitsAnAbsentName(t *testing.T) { w.WriteHeader(http.StatusNoContent) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - err := DeviceClaim([]string{"354820091234567", "3"}) + err := DeviceClaim(session, []string{"354820091234567", "3"}) if err != nil { t.Fatal(err) @@ -112,6 +112,10 @@ func TestDeviceClaimOmitsAnAbsentName(t *testing.T) { if nameWasPresent { t.Error("the request included a name although none was given") } + + if out.String() != "Press the button on the device to finish claiming it.\nClaimed the device into \"pilot\".\n" { + t.Errorf("output = %q", out.String()) + } } func TestDeviceClaimArguments(t *testing.T) { @@ -129,7 +133,7 @@ func TestDeviceClaimArguments(t *testing.T) { } for _, test := range tests { - err := DeviceClaim(test.arguments) + err := DeviceClaim(Session{}, test.arguments) if err == nil || !strings.Contains(err.Error(), test.wantError) { t.Errorf("%s: error = %v, want it to mention %q", test.name, err, test.wantError) @@ -143,9 +147,9 @@ func TestDeviceClaimUnknownFleetUsesFleetIdGuidance(t *testing.T) { fmt.Fprint(w, `[]`) }) - loggedInTestServer(t, mux) + session, _ := loggedInSession(t, mux) - err := DeviceClaim([]string{"354820091234567", "9"}) + err := DeviceClaim(session, []string{"354820091234567", "9"}) if err == nil || !strings.Contains(err.Error(), "shown by fleet list") { t.Fatalf("error = %v", err) diff --git a/internal/commands/device_list.go b/internal/commands/device_list.go index 3e6bb34..ccf09ba 100644 --- a/internal/commands/device_list.go +++ b/internal/commands/device_list.go @@ -4,23 +4,12 @@ import ( "encoding/json" "errors" "fmt" - "os" "strconv" "time" ) -func DeviceList(arguments []string) error { - jsonOutput := false - positionals := []string{} - - for _, argument := range arguments { - if argument == "--json" { - jsonOutput = true - continue - } - - positionals = append(positionals, argument) - } +func DeviceList(session Session, arguments []string) error { + positionals, jsonOutput := takeJsonFlag(arguments) if len(positionals) > 1 { return errors.New("device list takes at most one fleet id") @@ -38,13 +27,13 @@ func DeviceList(arguments []string) error { chosenFleetId = parsed } - devices, err := fetchDevices() + devices, err := fetchDevices(session) if err != nil { return err } - fleets, err := fetchFleets() + fleets, err := fetchFleets(session) if err != nil { return err @@ -71,14 +60,14 @@ func DeviceList(arguments []string) error { } if jsonOutput { - return json.NewEncoder(os.Stdout).Encode(filtered) + return json.NewEncoder(session.Out).Encode(filtered) } if len(filtered) == 0 { if chosenFleetId == 0 { - fmt.Println("No devices yet. Claim one with device claim.") + fmt.Fprintln(session.Out, "No devices yet. Claim one with device claim.") } else { - fmt.Println("No devices in that fleet.") + fmt.Fprintln(session.Out, "No devices in that fleet.") } return nil @@ -139,12 +128,12 @@ func DeviceList(arguments []string) error { storageWidth = max(storageWidth, len(storageValues[index])) } - fmt.Printf("%-*s %-*s %-*s %-*s %-*s %s\n", + fmt.Fprintf(session.Out, "%-*s %-*s %-*s %-*s %-*s %s\n", imeiWidth, "IMEI", nameWidth, "NAME", fleetWidth, "FLEET", stateWidth, "STATE", storageWidth, "STORAGE", "LAST SEEN") for index := range filtered { - fmt.Printf("%-*s %-*s %-*s %-*s %-*s %s\n", + fmt.Fprintf(session.Out, "%-*s %-*s %-*s %-*s %-*s %s\n", imeiWidth, imeiValues[index], nameWidth, nameValues[index], fleetWidth, fleetValues[index], stateWidth, stateValues[index], storageWidth, storageValues[index], lastSeenValues[index]) } diff --git a/internal/commands/device_list_test.go b/internal/commands/device_list_test.go index f603690..018f177 100644 --- a/internal/commands/device_list_test.go +++ b/internal/commands/device_list_test.go @@ -38,9 +38,11 @@ func TestDeviceList(t *testing.T) { mux := http.NewServeMux() mux.HandleFunc("GET /devices", func(w http.ResponseWriter, r *http.Request) { fmt.Fprint(w, devices) }) mux.HandleFunc("GET /fleets", func(w http.ResponseWriter, r *http.Request) { fmt.Fprint(w, fleets) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - printed, err := captureStdout(t, func() error { return DeviceList(test.arguments) }) + err := DeviceList(session, test.arguments) + + printed := out.String() if test.wantError != "" { if err == nil || !strings.Contains(err.Error(), test.wantError) { @@ -130,9 +132,11 @@ func TestDeviceListEmptyAndServerError(t *testing.T) { mux := http.NewServeMux() mux.HandleFunc("GET /devices", func(w http.ResponseWriter, r *http.Request) { fmt.Fprint(w, `[]`) }) mux.HandleFunc("GET /fleets", func(w http.ResponseWriter, r *http.Request) { fmt.Fprint(w, `[]`) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) + + err := DeviceList(session, nil) - printed, err := captureStdout(t, func() error { return DeviceList(nil) }) + printed := out.String() if err != nil { t.Fatal(err) @@ -146,9 +150,9 @@ func TestDeviceListEmptyAndServerError(t *testing.T) { errorMux.HandleFunc("GET /devices", func(w http.ResponseWriter, r *http.Request) { http.Error(w, "devices unavailable", http.StatusServiceUnavailable) }) - loggedInTestServer(t, errorMux) + errorSession, _ := loggedInSession(t, errorMux) - err = DeviceList(nil) + err = DeviceList(errorSession, nil) if err == nil || err.Error() != "the server said: devices unavailable" { t.Fatalf("error = %v", err) diff --git a/internal/commands/device_release.go b/internal/commands/device_release.go index a04065e..52824cb 100644 --- a/internal/commands/device_release.go +++ b/internal/commands/device_release.go @@ -4,13 +4,11 @@ import ( "bufio" "errors" "fmt" - "io" "net/http" - "os" "strings" ) -func DeviceRelease(arguments []string) error { +func DeviceRelease(session Session, arguments []string) error { if len(arguments) != 1 { return errors.New("device release takes an IMEI") } @@ -21,7 +19,7 @@ func DeviceRelease(arguments []string) error { return errors.New("the IMEI is the 15-digit number printed on the device") } - devices, err := fetchDevices() + devices, err := fetchDevices(session) if err != nil { return err @@ -39,7 +37,7 @@ func DeviceRelease(arguments []string) error { return errors.New("no such device, device list shows yours") } - fleets, err := fetchFleets() + fleets, err := fetchFleets(session) if err != nil { return err @@ -57,24 +55,24 @@ func DeviceRelease(arguments []string) error { return errors.New("no such device, device list shows yours") } - fmt.Printf("Release the device from %q? It erases everything on the device, and claiming it again means pressing its button in person. [y/N] ", fleetName) + fmt.Fprintf(session.Out, "Release the device from %q? It erases everything on the device, and claiming it again means pressing its button in person. [y/N] ", fleetName) - answer, _ := bufio.NewReader(os.Stdin).ReadString('\n') + answer, _ := bufio.NewReader(session.In).ReadString('\n') answer = strings.ToLower(strings.TrimSpace(answer)) if answer != "y" && answer != "yes" { - fmt.Println("Nothing released.") + fmt.Fprintln(session.Out, "Nothing released.") return nil } - request, err := authenticatedRequest(http.MethodDelete, "/devices/"+imei, nil) + request, err := authenticatedRequest(session, http.MethodDelete, "/devices/"+imei, nil) if err != nil { return err } - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return fmt.Errorf("the server could not be reached: %w", err) @@ -83,12 +81,10 @@ func DeviceRelease(arguments []string) error { defer response.Body.Close() if response.StatusCode != http.StatusNoContent { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return serverError(response) } - fmt.Println("Released the device.") + fmt.Fprintln(session.Out, "Released the device.") return nil } diff --git a/internal/commands/device_release_test.go b/internal/commands/device_release_test.go index 4d41310..92601aa 100644 --- a/internal/commands/device_release_test.go +++ b/internal/commands/device_release_test.go @@ -42,12 +42,12 @@ func TestDeviceRelease(t *testing.T) { w.WriteHeader(http.StatusNoContent) }) - loggedInTestServer(t, mux) - answerOnStdin(t, test.answer) + session, out := loggedInSession(t, mux) + session.In = strings.NewReader(test.answer) - printed, err := captureStdout(t, func() error { - return DeviceRelease([]string{"354820091234567"}) - }) + err := DeviceRelease(session, []string{"354820091234567"}) + + printed := out.String() if test.wantError != "" { if err == nil || err.Error() != test.wantError { @@ -85,7 +85,7 @@ func TestDeviceReleaseArgumentsAndUnknownDevice(t *testing.T) { } for _, test := range tests { - err := DeviceRelease(test.arguments) + err := DeviceRelease(Session{}, test.arguments) if err == nil || !strings.Contains(err.Error(), test.wantError) { t.Errorf("%s: error = %v", test.name, err) @@ -96,9 +96,9 @@ func TestDeviceReleaseArgumentsAndUnknownDevice(t *testing.T) { mux.HandleFunc("GET /devices", func(w http.ResponseWriter, r *http.Request) { fmt.Fprint(w, `[]`) }) - loggedInTestServer(t, mux) + session, _ := loggedInSession(t, mux) - err := DeviceRelease([]string{"354820091234567"}) + err := DeviceRelease(session, []string{"354820091234567"}) if err == nil || err.Error() != "no such device, device list shows yours" { t.Fatalf("error = %v", err) diff --git a/internal/commands/device_rename.go b/internal/commands/device_rename.go index dbb0441..55cf52b 100644 --- a/internal/commands/device_rename.go +++ b/internal/commands/device_rename.go @@ -5,12 +5,11 @@ import ( "encoding/json" "errors" "fmt" - "io" "net/http" "strings" ) -func DeviceRename(arguments []string) error { +func DeviceRename(session Session, arguments []string) error { if len(arguments) != 2 { return errors.New("device rename takes an IMEI and a name, quoted if it has spaces") } @@ -33,7 +32,7 @@ func DeviceRename(arguments []string) error { return err } - request, err := authenticatedRequest(http.MethodPatch, "/devices/"+imei, bytes.NewReader(body)) + request, err := authenticatedRequest(session, http.MethodPatch, "/devices/"+imei, bytes.NewReader(body)) if err != nil { return err @@ -41,7 +40,7 @@ func DeviceRename(arguments []string) error { request.Header.Set("Content-Type", "application/json") - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return fmt.Errorf("the server could not be reached: %w", err) @@ -50,12 +49,10 @@ func DeviceRename(arguments []string) error { defer response.Body.Close() if response.StatusCode != http.StatusNoContent { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return serverError(response) } - fmt.Printf("Renamed the device to %q.\n", name) + fmt.Fprintf(session.Out, "Renamed the device to %q.\n", name) return nil } diff --git a/internal/commands/device_rename_test.go b/internal/commands/device_rename_test.go index 184a9e6..aeea21e 100644 --- a/internal/commands/device_rename_test.go +++ b/internal/commands/device_rename_test.go @@ -23,11 +23,11 @@ func TestDeviceRename(t *testing.T) { w.WriteHeader(http.StatusNoContent) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - printed, err := captureStdout(t, func() error { - return DeviceRename([]string{"354820091234567", " pilot "}) - }) + err := DeviceRename(session, []string{"354820091234567", " pilot "}) + + printed := out.String() if err != nil { t.Fatal(err) @@ -48,9 +48,9 @@ func TestDeviceRenameServerError(t *testing.T) { http.Error(w, "no such device", http.StatusNotFound) }) - loggedInTestServer(t, mux) + session, _ := loggedInSession(t, mux) - err := DeviceRename([]string{"354820091234567", "pilot"}) + err := DeviceRename(session, []string{"354820091234567", "pilot"}) if err == nil || err.Error() != "the server said: no such device" { t.Fatalf("error = %v", err) @@ -73,7 +73,7 @@ func TestDeviceRenameArguments(t *testing.T) { } for _, test := range tests { - err := DeviceRename(test.arguments) + err := DeviceRename(Session{}, test.arguments) if err == nil || !strings.Contains(err.Error(), test.wantError) { t.Errorf("%s: error = %v, want it to mention %q", test.name, err, test.wantError) diff --git a/internal/commands/devices.go b/internal/commands/devices.go index acbf7a7..f7ed7f1 100644 --- a/internal/commands/devices.go +++ b/internal/commands/devices.go @@ -3,9 +3,7 @@ package commands import ( "encoding/json" "fmt" - "io" "net/http" - "strings" ) type deviceEntry struct { @@ -18,14 +16,14 @@ type deviceEntry struct { StorageTotal *int64 `json:"storage_total"` } -func fetchDevices() ([]deviceEntry, error) { - request, err := authenticatedRequest(http.MethodGet, "/devices", nil) +func fetchDevices(session Session) ([]deviceEntry, error) { + request, err := authenticatedRequest(session, http.MethodGet, "/devices", nil) if err != nil { return nil, err } - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return nil, fmt.Errorf("the server could not be reached: %w", err) @@ -34,9 +32,7 @@ func fetchDevices() ([]deviceEntry, error) { defer response.Body.Close() if response.StatusCode != http.StatusOK { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return nil, fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return nil, serverError(response) } devices := []deviceEntry{} diff --git a/internal/commands/dispatch.go b/internal/commands/dispatch.go new file mode 100644 index 0000000..74b26a6 --- /dev/null +++ b/internal/commands/dispatch.go @@ -0,0 +1,163 @@ +package commands + +import ( + "fmt" + "io" + "strings" +) + +type Command struct { + Name string + Arguments string + Summary string + Run func(session Session, arguments []string) error +} + +type Section struct { + Title string + Commands []Command +} + +func resolve(sections []Section, arguments []string) (Command, []string, bool) { + longest := Command{} + longestWords := 0 + + for _, section := range sections { + for _, candidate := range section.Commands { + words := strings.Fields(candidate.Name) + + if len(words) > len(arguments) || len(words) <= longestWords { + continue + } + + matches := true + + for index, word := range words { + if arguments[index] != word { + matches = false + break + } + } + + if !matches { + continue + } + + longest = candidate + longestWords = len(words) + } + } + + if longestWords == 0 { + return Command{}, nil, false + } + + return longest, arguments[longestWords:], true +} + +func printHelp(session Session, sections []Section) { + widest := 0 + + for _, section := range sections { + for _, entry := range section.Commands { + width := len(entry.Name) + + if entry.Arguments != "" { + width += 1 + len(entry.Arguments) + } + + if width > widest { + widest = width + } + } + } + + fmt.Fprintf(session.Out, "superstack %s\n\n", session.Version) + fmt.Fprint(session.Out, "Usage: superstack [arguments]\n") + + for _, section := range sections { + fmt.Fprintf(session.Out, "\n%s\n", section.Title) + + for _, entry := range section.Commands { + signature := entry.Name + + if entry.Arguments != "" { + signature += " " + entry.Arguments + } + + fmt.Fprintf(session.Out, " %-*s %s\n", widest, signature, entry.Summary) + } + } +} + +func Dispatch(sections []Section, version string, arguments []string, in io.Reader, out io.Writer) error { + arguments, base, err := TakeServerFlag(arguments) + + if err != nil { + return err + } + + session := NewSession(base, version, in, out) + + if len(arguments) == 0 { + printHelp(session, sections) + return nil + } + + switch arguments[0] { + case "-h", "--help": + printHelp(session, sections) + return nil + + case "-v", "--version": + fmt.Fprintln(session.Out, session.Version) + return nil + } + + entry, rest, found := resolve(sections, arguments) + + if !found { + return fmt.Errorf("unknown command %q\nRun 'superstack help' for the list.", strings.Join(arguments, " ")) + } + + switch entry.Name { + case "version": + fmt.Fprintln(session.Out, session.Version) + return nil + + case "help": + if len(rest) == 0 { + printHelp(session, sections) + return nil + } + + topic, _, topicFound := resolve(sections, rest) + + if !topicFound { + return fmt.Errorf("unknown command %q", strings.Join(rest, " ")) + } + + signature := topic.Name + + if topic.Arguments != "" { + signature += " " + topic.Arguments + } + + fmt.Fprintf(session.Out, "superstack %s\n\n %s\n", signature, topic.Summary) + return nil + } + + if entry.Run == nil { + return fmt.Errorf("%s is not implemented yet", entry.Name) + } + + err = CheckServer(session) + + if err != nil { + return err + } + + err = entry.Run(session, rest) + + return err +} diff --git a/internal/commands/dispatch_test.go b/internal/commands/dispatch_test.go new file mode 100644 index 0000000..cdf1713 --- /dev/null +++ b/internal/commands/dispatch_test.go @@ -0,0 +1,94 @@ +package commands + +import ( + "bytes" + "slices" + "strings" + "testing" +) + +func TestResolve(t *testing.T) { + sections := []Section{{Commands: []Command{ + {Name: "login"}, + {Name: "device list"}, + {Name: "device claim"}, + {Name: "fleet create"}, + {Name: "member add"}, + {Name: "key create"}, + {Name: "account balance"}, + {Name: "account topup"}, + {Name: "upload"}, + }}} + + tests := []struct { + arguments []string + name string + rest []string + found bool + }{ + {arguments: []string{"login"}, name: "login", rest: []string{}, found: true}, + {arguments: []string{"device", "list"}, name: "device list", rest: []string{}, found: true}, + {arguments: []string{"device", "claim", "354820091234567", "sensor-01"}, name: "device claim", rest: []string{"354820091234567", "sensor-01"}, found: true}, + {arguments: []string{"fleet", "create", "thermostats"}, name: "fleet create", rest: []string{"thermostats"}, found: true}, + {arguments: []string{"member", "add", "member@example.com"}, name: "member add", rest: []string{"member@example.com"}, found: true}, + {arguments: []string{"key", "create", "42", "production"}, name: "key create", rest: []string{"42", "production"}, found: true}, + {arguments: []string{"account", "balance"}, name: "account balance", rest: []string{}, found: true}, + {arguments: []string{"account", "topup", "42"}, name: "account topup", rest: []string{"42"}, found: true}, + {arguments: []string{"upload", "./main.lua", "--device", "sensor-01"}, name: "upload", rest: []string{"./main.lua", "--device", "sensor-01"}, found: true}, + {arguments: []string{"fleet"}, found: false}, + {arguments: []string{"member"}, found: false}, + {arguments: []string{"device"}, found: false}, + {arguments: []string{"key"}, found: false}, + {arguments: []string{"account"}, found: false}, + {arguments: []string{"deploy"}, found: false}, + {arguments: []string{}, found: false}, + } + + for _, test := range tests { + entry, rest, found := resolve(sections, test.arguments) + + if found != test.found { + t.Errorf("resolve(%q) found = %v, want %v", test.arguments, found, test.found) + continue + } + + if !found { + continue + } + + if entry.Name != test.name { + t.Errorf("resolve(%q) name = %q, want %q", test.arguments, entry.Name, test.name) + } + + if !slices.Equal(rest, test.rest) { + t.Errorf("resolve(%q) rest = %q, want %q", test.arguments, rest, test.rest) + } + } +} + +func TestHelpListsEveryCommand(t *testing.T) { + sections := []Section{ + {Title: "Things", Commands: []Command{{Name: "thing list", Arguments: "[--json]", Summary: "List things"}}}, + {Title: "Account", Commands: []Command{{Name: "account delete", Summary: "Delete the account"}}}, + } + out := &bytes.Buffer{} + session := NewSession(defaultApiBase, "1.2.3", strings.NewReader(""), out) + + printHelp(session, sections) + + for _, section := range sections { + if !strings.Contains(out.String(), section.Title) { + t.Errorf("help is missing the section %q", section.Title) + } + + for _, entry := range section.Commands { + if !strings.Contains(out.String(), entry.Name) { + t.Errorf("help is missing the command %q", entry.Name) + } + + if !strings.Contains(out.String(), entry.Summary) { + t.Errorf("help is missing the summary for %q", entry.Name) + } + } + } +} diff --git a/internal/commands/fleet_create.go b/internal/commands/fleet_create.go index e8a8f41..afa2780 100644 --- a/internal/commands/fleet_create.go +++ b/internal/commands/fleet_create.go @@ -5,12 +5,10 @@ import ( "encoding/json" "errors" "fmt" - "io" "net/http" - "strings" ) -func FleetCreate(arguments []string) error { +func FleetCreate(session Session, arguments []string) error { if len(arguments) != 1 || arguments[0] == "" { return errors.New("fleet create takes one name, quoted if it has spaces") @@ -22,7 +20,7 @@ func FleetCreate(arguments []string) error { return err } - request, err := authenticatedRequest(http.MethodPost, "/fleets", bytes.NewReader(body)) + request, err := authenticatedRequest(session, http.MethodPost, "/fleets", bytes.NewReader(body)) if err != nil { return err @@ -30,7 +28,7 @@ func FleetCreate(arguments []string) error { request.Header.Set("Content-Type", "application/json") - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return fmt.Errorf("the server could not be reached: %w", err) @@ -39,9 +37,7 @@ func FleetCreate(arguments []string) error { defer response.Body.Close() if response.StatusCode != http.StatusOK { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return serverError(response) } created := struct { @@ -55,7 +51,7 @@ func FleetCreate(arguments []string) error { return err } - fmt.Printf("Created fleet %q with id %d.\n", created.Name, created.Id) + fmt.Fprintf(session.Out, "Created fleet %q with id %d.\n", created.Name, created.Id) return nil } diff --git a/internal/commands/fleet_create_test.go b/internal/commands/fleet_create_test.go index 04c9b45..85bec25 100644 --- a/internal/commands/fleet_create_test.go +++ b/internal/commands/fleet_create_test.go @@ -25,9 +25,9 @@ func TestFleetCreate(t *testing.T) { fmt.Fprintf(w, `{"id": 5, "name": %q}`, body.Name) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - err := FleetCreate([]string{"field trial"}) + err := FleetCreate(session, []string{"field trial"}) if err != nil { t.Fatal(err) @@ -36,6 +36,10 @@ func TestFleetCreate(t *testing.T) { if created != "field trial" { t.Errorf("the server saw %q created, want %q", created, "field trial") } + + if out.String() != "Created fleet \"field trial\" with id 5.\n" { + t.Errorf("output = %q", out.String()) + } } func TestFleetCreateRelaysARefusal(t *testing.T) { @@ -45,9 +49,9 @@ func TestFleetCreateRelaysARefusal(t *testing.T) { http.Error(w, "the body must carry a name", http.StatusBadRequest) }) - loggedInTestServer(t, mux) + session, _ := loggedInSession(t, mux) - err := FleetCreate([]string{"field trial"}) + err := FleetCreate(session, []string{"field trial"}) if err == nil || !strings.Contains(err.Error(), "the server said: the body must carry a name") { t.Fatalf("error = %v, want the relayed refusal", err) @@ -65,7 +69,7 @@ func TestFleetCreateTakesOneName(t *testing.T) { } for _, test := range tests { - err := FleetCreate(test.arguments) + err := FleetCreate(Session{}, test.arguments) if err == nil || !strings.Contains(err.Error(), "takes one name") { t.Errorf("%s: error = %v, want the one-name hint", test.name, err) diff --git a/internal/commands/fleet_delete.go b/internal/commands/fleet_delete.go index 222e2d8..b105f48 100644 --- a/internal/commands/fleet_delete.go +++ b/internal/commands/fleet_delete.go @@ -4,14 +4,12 @@ import ( "bufio" "errors" "fmt" - "io" "net/http" - "os" "strconv" "strings" ) -func FleetDelete(arguments []string) error { +func FleetDelete(session Session, arguments []string) error { if len(arguments) != 1 { return errors.New("fleet delete takes a fleet id") @@ -23,7 +21,7 @@ func FleetDelete(arguments []string) error { return errors.New("the fleet id is the number shown by fleet list") } - fleets, err := fetchFleets() + fleets, err := fetchFleets(session) if err != nil { return err @@ -43,7 +41,7 @@ func FleetDelete(arguments []string) error { return errors.New("no such fleet") } - balances, err := fetchBalances() + balances, err := fetchBalances(session) if err != nil { return err @@ -66,28 +64,28 @@ func FleetDelete(arguments []string) error { consequence := "It erases them all, and claiming one again means pressing its button in person." if forfeited == "" { - fmt.Printf("Delete %q and release its devices? %s [y/N] ", name, consequence) + fmt.Fprintf(session.Out, "Delete %q and release its devices? %s [y/N] ", name, consequence) } else { - fmt.Printf("Delete %q, release its devices, and forfeit its remaining %s of credit? %s [y/N] ", name, forfeited, consequence) + fmt.Fprintf(session.Out, "Delete %q, release its devices, and forfeit its remaining %s of credit? %s [y/N] ", name, forfeited, consequence) } - answer, _ := bufio.NewReader(os.Stdin).ReadString('\n') + answer, _ := bufio.NewReader(session.In).ReadString('\n') answer = strings.ToLower(strings.TrimSpace(answer)) if answer != "y" && answer != "yes" { - fmt.Println("Nothing deleted.") + fmt.Fprintln(session.Out, "Nothing deleted.") return nil } - request, err := authenticatedRequest(http.MethodDelete, + request, err := authenticatedRequest(session, http.MethodDelete, "/fleets/"+strconv.FormatInt(fleetId, 10), nil) if err != nil { return err } - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return fmt.Errorf("the server could not be reached: %w", err) @@ -96,12 +94,10 @@ func FleetDelete(arguments []string) error { defer response.Body.Close() if response.StatusCode != http.StatusNoContent { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return serverError(response) } - fmt.Printf("Deleted %q.\n", name) + fmt.Fprintf(session.Out, "Deleted %q.\n", name) return nil } diff --git a/internal/commands/fleet_delete_test.go b/internal/commands/fleet_delete_test.go index 223684c..c99bc75 100644 --- a/internal/commands/fleet_delete_test.go +++ b/internal/commands/fleet_delete_test.go @@ -3,37 +3,10 @@ package commands import ( "fmt" "net/http" - "os" "strings" "testing" ) -func answerOnStdin(t *testing.T, answer string) { - t.Helper() - - readEnd, writeEnd, err := os.Pipe() - - if err != nil { - t.Fatal(err) - } - - originalStdin := os.Stdin - - os.Stdin = readEnd - - t.Cleanup(func() { os.Stdin = originalStdin }) - - if answer != "" { - _, err = writeEnd.WriteString(answer) - - if err != nil { - t.Fatal(err) - } - } - - writeEnd.Close() -} - func TestFleetDelete(t *testing.T) { tests := []struct { name string @@ -81,13 +54,13 @@ func TestFleetDelete(t *testing.T) { w.WriteHeader(http.StatusNoContent) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - answerOnStdin(t, test.answer) + session.In = strings.NewReader(test.answer) - printed, err := captureStdout(t, func() error { - return FleetDelete([]string{"3"}) - }) + err := FleetDelete(session, []string{"3"}) + + printed := out.String() switch { case test.wantError != "": @@ -146,13 +119,13 @@ func TestFleetDeletePromptStatesForfeitedCredit(t *testing.T) { fmt.Fprint(w, test.balance) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - answerOnStdin(t, "n\n") + session.In = strings.NewReader("n\n") - printed, err := captureStdout(t, func() error { - return FleetDelete([]string{"3"}) - }) + err := FleetDelete(session, []string{"3"}) + + printed := out.String() if err != nil { t.Fatal(err) @@ -193,11 +166,11 @@ func TestFleetDeleteRefusesWhenTheBalanceIsUnknown(t *testing.T) { w.WriteHeader(http.StatusNoContent) }) - loggedInTestServer(t, mux) + session, _ := loggedInSession(t, mux) - answerOnStdin(t, "y\n") + session.In = strings.NewReader("y\n") - err := FleetDelete([]string{"3"}) + err := FleetDelete(session, []string{"3"}) if err == nil || !strings.Contains(err.Error(), "could not read the balances") { t.Fatalf("error = %v, want the server's balance refusal", err) @@ -215,9 +188,9 @@ func TestFleetDeleteUnknownFleet(t *testing.T) { fmt.Fprint(w, `[]`) }) - loggedInTestServer(t, mux) + session, _ := loggedInSession(t, mux) - err := FleetDelete([]string{"9"}) + err := FleetDelete(session, []string{"9"}) if err == nil || !strings.Contains(err.Error(), "no such fleet") { t.Fatalf("error = %v, want no such fleet", err) @@ -236,7 +209,7 @@ func TestFleetDeleteArguments(t *testing.T) { } for _, test := range tests { - err := FleetDelete(test.arguments) + err := FleetDelete(Session{}, test.arguments) if err == nil || !strings.Contains(err.Error(), test.wantError) { t.Errorf("%s: error = %v, want it to mention %q", test.name, err, test.wantError) diff --git a/internal/commands/fleet_list.go b/internal/commands/fleet_list.go index ee12986..95faf23 100644 --- a/internal/commands/fleet_list.go +++ b/internal/commands/fleet_list.go @@ -3,34 +3,29 @@ package commands import ( "encoding/json" "fmt" - "os" "strconv" ) -func FleetList(arguments []string) error { +func FleetList(session Session, arguments []string) error { - jsonOutput := false + positionals, jsonOutput := takeJsonFlag(arguments) - for _, argument := range arguments { - if argument != "--json" { - return fmt.Errorf("fleet list takes no arguments, only --json") - } - - jsonOutput = true + if len(positionals) != 0 { + return fmt.Errorf("fleet list takes no arguments, only --json") } - fleets, err := fetchFleets() + fleets, err := fetchFleets(session) if err != nil { return err } if jsonOutput { - return json.NewEncoder(os.Stdout).Encode(fleets) + return json.NewEncoder(session.Out).Encode(fleets) } if len(fleets) == 0 { - fmt.Println("No fleets yet. Create one with fleet create.") + fmt.Fprintln(session.Out, "No fleets yet. Create one with fleet create.") return nil } @@ -42,7 +37,7 @@ func FleetList(arguments []string) error { nameWidth = max(nameWidth, len(fleet.Name)) } - fmt.Printf("%-*s %-*s %s\n", idWidth, "ID", nameWidth, "NAME", "ROLE") + fmt.Fprintf(session.Out, "%-*s %-*s %s\n", idWidth, "ID", nameWidth, "NAME", "ROLE") for _, fleet := range fleets { role := "member" @@ -51,7 +46,7 @@ func FleetList(arguments []string) error { role = "owner" } - fmt.Printf("%-*d %-*s %s\n", idWidth, fleet.Id, nameWidth, fleet.Name, role) + fmt.Fprintf(session.Out, "%-*d %-*s %s\n", idWidth, fleet.Id, nameWidth, fleet.Name, role) } return nil diff --git a/internal/commands/fleet_list_test.go b/internal/commands/fleet_list_test.go index 2091f2f..0cdc6d4 100644 --- a/internal/commands/fleet_list_test.go +++ b/internal/commands/fleet_list_test.go @@ -49,11 +49,11 @@ func TestFleetList(t *testing.T) { fmt.Fprint(w, test.fleets) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - printed, err := captureStdout(t, func() error { - return FleetList(test.arguments) - }) + err := FleetList(session, test.arguments) + + printed := out.String() if test.wantError != "" { if err == nil || !strings.Contains(err.Error(), test.wantError) { diff --git a/internal/commands/fleet_rename.go b/internal/commands/fleet_rename.go index 65357f2..8274467 100644 --- a/internal/commands/fleet_rename.go +++ b/internal/commands/fleet_rename.go @@ -5,13 +5,12 @@ import ( "encoding/json" "errors" "fmt" - "io" "net/http" "strconv" "strings" ) -func FleetRename(arguments []string) error { +func FleetRename(session Session, arguments []string) error { if len(arguments) != 2 { return errors.New("fleet rename takes a fleet id and a name, quoted if it has spaces") @@ -35,7 +34,7 @@ func FleetRename(arguments []string) error { return err } - request, err := authenticatedRequest(http.MethodPatch, + request, err := authenticatedRequest(session, http.MethodPatch, "/fleets/"+strconv.FormatInt(fleetId, 10), bytes.NewReader(body)) if err != nil { @@ -44,7 +43,7 @@ func FleetRename(arguments []string) error { request.Header.Set("Content-Type", "application/json") - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return fmt.Errorf("the server could not be reached: %w", err) @@ -53,12 +52,10 @@ func FleetRename(arguments []string) error { defer response.Body.Close() if response.StatusCode != http.StatusNoContent { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return serverError(response) } - fmt.Printf("Renamed the fleet to %q.\n", name) + fmt.Fprintf(session.Out, "Renamed the fleet to %q.\n", name) return nil } diff --git a/internal/commands/fleet_rename_test.go b/internal/commands/fleet_rename_test.go index 6399e9e..d861ba3 100644 --- a/internal/commands/fleet_rename_test.go +++ b/internal/commands/fleet_rename_test.go @@ -26,9 +26,9 @@ func TestFleetRename(t *testing.T) { w.WriteHeader(http.StatusNoContent) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - err := FleetRename([]string{"9", " pilot "}) + err := FleetRename(session, []string{"9", " pilot "}) if err != nil { t.Fatal(err) @@ -38,6 +38,10 @@ func TestFleetRename(t *testing.T) { t.Errorf("the server saw %q renamed to %q, want %q renamed to %q", renamedPath, renamedTo, "/fleets/9", "pilot") } + + if out.String() != "Renamed the fleet to \"pilot\".\n" { + t.Errorf("output = %q", out.String()) + } } func TestFleetRenameArguments(t *testing.T) { @@ -55,7 +59,7 @@ func TestFleetRenameArguments(t *testing.T) { } for _, test := range tests { - err := FleetRename(test.arguments) + err := FleetRename(Session{}, test.arguments) if err == nil || !strings.Contains(err.Error(), test.wantError) { t.Errorf("%s: error = %v, want it to mention %q", test.name, err, test.wantError) diff --git a/internal/commands/fleet_transfer.go b/internal/commands/fleet_transfer.go index 94ce714..6a12439 100644 --- a/internal/commands/fleet_transfer.go +++ b/internal/commands/fleet_transfer.go @@ -5,13 +5,11 @@ import ( "encoding/json" "errors" "fmt" - "io" "net/http" "strconv" - "strings" ) -func FleetTransfer(arguments []string) error { +func FleetTransfer(session Session, arguments []string) error { if len(arguments) != 2 || arguments[1] == "" { return errors.New("fleet transfer takes a fleet id and an email address") @@ -31,7 +29,7 @@ func FleetTransfer(arguments []string) error { return err } - request, err := authenticatedRequest(http.MethodPost, + request, err := authenticatedRequest(session, http.MethodPost, "/fleets/"+strconv.FormatInt(fleetId, 10)+"/owner", bytes.NewReader(body)) if err != nil { @@ -40,7 +38,7 @@ func FleetTransfer(arguments []string) error { request.Header.Set("Content-Type", "application/json") - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return fmt.Errorf("the server could not be reached: %w", err) @@ -49,12 +47,10 @@ func FleetTransfer(arguments []string) error { defer response.Body.Close() if response.StatusCode != http.StatusNoContent { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return serverError(response) } - fmt.Printf("Transferred the fleet to %s.\n", email) + fmt.Fprintf(session.Out, "Transferred the fleet to %s.\n", email) return nil } diff --git a/internal/commands/fleet_transfer_test.go b/internal/commands/fleet_transfer_test.go index fdcbed1..e03fe78 100644 --- a/internal/commands/fleet_transfer_test.go +++ b/internal/commands/fleet_transfer_test.go @@ -26,9 +26,9 @@ func TestFleetTransfer(t *testing.T) { w.WriteHeader(http.StatusNoContent) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - err := FleetTransfer([]string{"3", "successor@example.com"}) + err := FleetTransfer(session, []string{"3", "successor@example.com"}) if err != nil { t.Fatal(err) @@ -38,6 +38,10 @@ func TestFleetTransfer(t *testing.T) { t.Errorf("the server saw %q handed to %q, want %q handed to %q", transferredPath, transferredTo, "/fleets/3/owner", "successor@example.com") } + + if out.String() != "Transferred the fleet to successor@example.com.\n" { + t.Errorf("output = %q", out.String()) + } } func TestFleetTransferArguments(t *testing.T) { @@ -53,7 +57,7 @@ func TestFleetTransferArguments(t *testing.T) { } for _, test := range tests { - err := FleetTransfer(test.arguments) + err := FleetTransfer(Session{}, test.arguments) if err == nil || !strings.Contains(err.Error(), test.wantError) { t.Errorf("%s: error = %v, want it to mention %q", test.name, err, test.wantError) diff --git a/internal/commands/fleets.go b/internal/commands/fleets.go index 6179c75..b7b6e48 100644 --- a/internal/commands/fleets.go +++ b/internal/commands/fleets.go @@ -3,9 +3,7 @@ package commands import ( "encoding/json" "fmt" - "io" "net/http" - "strings" ) type fleetEntry struct { @@ -14,14 +12,14 @@ type fleetEntry struct { Owner bool `json:"owner"` } -func fetchFleets() ([]fleetEntry, error) { - request, err := authenticatedRequest(http.MethodGet, "/fleets", nil) +func fetchFleets(session Session) ([]fleetEntry, error) { + request, err := authenticatedRequest(session, http.MethodGet, "/fleets", nil) if err != nil { return nil, err } - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return nil, fmt.Errorf("the server could not be reached: %w", err) @@ -30,9 +28,7 @@ func fetchFleets() ([]fleetEntry, error) { defer response.Body.Close() if response.StatusCode != http.StatusOK { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return nil, fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return nil, serverError(response) } fleets := []fleetEntry{} diff --git a/internal/commands/fleets_test.go b/internal/commands/fleets_test.go index 7da86d2..74901f9 100644 --- a/internal/commands/fleets_test.go +++ b/internal/commands/fleets_test.go @@ -10,7 +10,7 @@ import ( func TestFetchFleetsNotLoggedIn(t *testing.T) { isolateKeyStorage(t) - _, err := fetchFleets() + _, err := fetchFleets(Session{}) if err == nil || !strings.Contains(err.Error(), "not logged in") { t.Fatalf("error = %v, want the not-logged-in hint", err) @@ -18,7 +18,7 @@ func TestFetchFleetsNotLoggedIn(t *testing.T) { } func TestFetchFleetsEmptyKeyFile(t *testing.T) { - loggedInTestServer(t, http.NotFoundHandler()) + session, _ := loggedInSession(t, http.NotFoundHandler()) path, err := keyPath() @@ -32,7 +32,7 @@ func TestFetchFleetsEmptyKeyFile(t *testing.T) { t.Fatal(err) } - _, err = fetchFleets() + _, err = fetchFleets(session) if err == nil || !strings.Contains(err.Error(), "not logged in") { t.Fatalf("error = %v, want the not-logged-in hint for an empty key file", err) diff --git a/internal/commands/key_create.go b/internal/commands/key_create.go index 7131c7b..7d5358a 100644 --- a/internal/commands/key_create.go +++ b/internal/commands/key_create.go @@ -5,13 +5,11 @@ import ( "encoding/json" "errors" "fmt" - "io" "net/http" "strconv" - "strings" ) -func KeyCreate(arguments []string) error { +func KeyCreate(session Session, arguments []string) error { if len(arguments) != 2 || arguments[1] == "" { return errors.New("key create takes a fleet id and a label, quoted if it has spaces") @@ -29,7 +27,7 @@ func KeyCreate(arguments []string) error { return err } - request, err := authenticatedRequest(http.MethodPost, + request, err := authenticatedRequest(session, http.MethodPost, "/fleets/"+strconv.FormatInt(fleetId, 10)+"/keys", bytes.NewReader(body)) if err != nil { @@ -38,7 +36,7 @@ func KeyCreate(arguments []string) error { request.Header.Set("Content-Type", "application/json") - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return fmt.Errorf("the server could not be reached: %w", err) @@ -47,9 +45,7 @@ func KeyCreate(arguments []string) error { defer response.Body.Close() if response.StatusCode != http.StatusOK { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return serverError(response) } created := struct { @@ -63,7 +59,7 @@ func KeyCreate(arguments []string) error { return err } - fmt.Printf("Created key %d.\n\n %s\n\nAnyone holding it can send data to the fleet, and it is shown only this once.\n", created.Id, created.Key) + fmt.Fprintf(session.Out, "Created key %d.\n\n %s\n\nAnyone holding it can send data to the fleet, and it is shown only this once.\n", created.Id, created.Key) return nil } diff --git a/internal/commands/key_create_test.go b/internal/commands/key_create_test.go index 2dd82ac..93a8335 100644 --- a/internal/commands/key_create_test.go +++ b/internal/commands/key_create_test.go @@ -88,11 +88,11 @@ func TestKeyCreate(t *testing.T) { fmt.Fprint(w, `{"id":1,"key":"ssf_testtesttestab2de"}`) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - printed, err := captureStdout(t, func() error { - return KeyCreate(test.arguments) - }) + err := KeyCreate(session, test.arguments) + + printed := out.String() if test.wantError != "" { if err == nil || !strings.Contains(err.Error(), test.wantError) { diff --git a/internal/commands/key_list.go b/internal/commands/key_list.go index 295dd30..58e5ada 100644 --- a/internal/commands/key_list.go +++ b/internal/commands/key_list.go @@ -4,24 +4,12 @@ import ( "encoding/json" "errors" "fmt" - "os" "strconv" ) -func KeyList(arguments []string) error { +func KeyList(session Session, arguments []string) error { - jsonOutput := false - - positionals := []string{} - - for _, argument := range arguments { - if argument == "--json" { - jsonOutput = true - continue - } - - positionals = append(positionals, argument) - } + positionals, jsonOutput := takeJsonFlag(arguments) if len(positionals) > 1 { return errors.New("key list takes at most one fleet id") @@ -39,7 +27,7 @@ func KeyList(arguments []string) error { chosenFleetId = parsed } - fleets, err := fetchFleets() + fleets, err := fetchFleets(session) if err != nil { return err @@ -57,7 +45,7 @@ func KeyList(arguments []string) error { } } - fetched, err := fetchKeys() + fetched, err := fetchKeys(session) if err != nil { return err @@ -72,11 +60,11 @@ func KeyList(arguments []string) error { } if jsonOutput { - return json.NewEncoder(os.Stdout).Encode(keys) + return json.NewEncoder(session.Out).Encode(keys) } if len(keys) == 0 { - fmt.Println("No keys yet. Create one with key create.") + fmt.Fprintln(session.Out, "No keys yet. Create one with key create.") return nil } @@ -90,11 +78,11 @@ func KeyList(arguments []string) error { fleetNameWidth = max(fleetNameWidth, len(fleetNames[key.Fleet])) } - fmt.Printf("%-*s %-*s %-*s %-8s %s\n", + fmt.Fprintf(session.Out, "%-*s %-*s %-*s %-8s %s\n", idWidth, "ID", fleetIdWidth, "FLEET", fleetNameWidth, "FLEET NAME", "KEY", "LABEL") for _, key := range keys { - fmt.Printf("%-*d %-*d %-*s ...%s %s\n", + fmt.Fprintf(session.Out, "%-*d %-*d %-*s ...%s %s\n", idWidth, key.Id, fleetIdWidth, key.Fleet, fleetNameWidth, fleetNames[key.Fleet], key.Suffix, key.Label) } diff --git a/internal/commands/key_list_test.go b/internal/commands/key_list_test.go index 709e485..3d8bf11 100644 --- a/internal/commands/key_list_test.go +++ b/internal/commands/key_list_test.go @@ -80,11 +80,11 @@ func TestKeyList(t *testing.T) { fmt.Fprint(w, keys) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - printed, err := captureStdout(t, func() error { - return KeyList(test.arguments) - }) + err := KeyList(session, test.arguments) + + printed := out.String() if test.wantError != "" { if err == nil || !strings.Contains(err.Error(), test.wantError) { diff --git a/internal/commands/key_revoke.go b/internal/commands/key_revoke.go index edf4364..890c296 100644 --- a/internal/commands/key_revoke.go +++ b/internal/commands/key_revoke.go @@ -4,14 +4,12 @@ import ( "bufio" "errors" "fmt" - "io" "net/http" - "os" "strconv" "strings" ) -func KeyRevoke(arguments []string) error { +func KeyRevoke(session Session, arguments []string) error { if len(arguments) != 1 { return errors.New("key revoke takes a key id") @@ -23,7 +21,7 @@ func KeyRevoke(arguments []string) error { return errors.New("the key id is the number shown by key list") } - keys, err := fetchKeys() + keys, err := fetchKeys(session) if err != nil { return err @@ -43,25 +41,25 @@ func KeyRevoke(arguments []string) error { return errors.New("no such key") } - fmt.Printf("Revoke %q? Anything still using it stops reaching the fleet. [y/N] ", label) + fmt.Fprintf(session.Out, "Revoke %q? Anything still using it stops reaching the fleet. [y/N] ", label) - answer, _ := bufio.NewReader(os.Stdin).ReadString('\n') + answer, _ := bufio.NewReader(session.In).ReadString('\n') answer = strings.ToLower(strings.TrimSpace(answer)) if answer != "y" && answer != "yes" { - fmt.Println("Nothing revoked.") + fmt.Fprintln(session.Out, "Nothing revoked.") return nil } - request, err := authenticatedRequest(http.MethodDelete, + request, err := authenticatedRequest(session, http.MethodDelete, "/keys/"+strconv.FormatInt(keyId, 10), nil) if err != nil { return err } - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return fmt.Errorf("the server could not be reached: %w", err) @@ -70,12 +68,10 @@ func KeyRevoke(arguments []string) error { defer response.Body.Close() if response.StatusCode != http.StatusNoContent { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return serverError(response) } - fmt.Printf("Revoked key %d.\n", keyId) + fmt.Fprintf(session.Out, "Revoked key %d.\n", keyId) return nil } diff --git a/internal/commands/key_revoke_test.go b/internal/commands/key_revoke_test.go index 37ac2f3..f394a7f 100644 --- a/internal/commands/key_revoke_test.go +++ b/internal/commands/key_revoke_test.go @@ -93,13 +93,12 @@ func TestKeyRevoke(t *testing.T) { w.WriteHeader(http.StatusNoContent) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) + session.In = strings.NewReader(test.answer) - answerOnStdin(t, test.answer) + err := KeyRevoke(session, test.arguments) - printed, err := captureStdout(t, func() error { - return KeyRevoke(test.arguments) - }) + printed := out.String() if test.wantError != "" { if err == nil || !strings.Contains(err.Error(), test.wantError) { diff --git a/internal/commands/keys.go b/internal/commands/keys.go index 6b64800..d8e132c 100644 --- a/internal/commands/keys.go +++ b/internal/commands/keys.go @@ -3,9 +3,7 @@ package commands import ( "encoding/json" "fmt" - "io" "net/http" - "strings" ) type keyEntry struct { @@ -15,14 +13,14 @@ type keyEntry struct { Suffix string `json:"suffix"` } -func fetchKeys() ([]keyEntry, error) { - request, err := authenticatedRequest(http.MethodGet, "/keys", nil) +func fetchKeys(session Session) ([]keyEntry, error) { + request, err := authenticatedRequest(session, http.MethodGet, "/keys", nil) if err != nil { return nil, err } - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return nil, fmt.Errorf("the server could not be reached: %w", err) @@ -31,9 +29,7 @@ func fetchKeys() ([]keyEntry, error) { defer response.Body.Close() if response.StatusCode != http.StatusOK { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return nil, fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return nil, serverError(response) } keys := []keyEntry{} diff --git a/internal/commands/login.go b/internal/commands/login.go index e206872..b438089 100644 --- a/internal/commands/login.go +++ b/internal/commands/login.go @@ -6,7 +6,6 @@ import ( "encoding/json" "errors" "fmt" - "io" "net/http" "net/url" "os" @@ -15,14 +14,8 @@ import ( "time" ) -var githubBase = "https://github.com" -var gitlabBase = "https://gitlab.com" - -var oauthClient = &http.Client{Timeout: 30 * time.Second} - -var minimumPollInterval = 5 - -func Login(arguments []string) error { +func Login(session Session, arguments []string) error { + oauthClient := &http.Client{Timeout: 30 * time.Second} if len(arguments) != 1 || (arguments[0] != "github" && arguments[0] != "gitlab") { return errors.New("login takes a provider: github or gitlab") @@ -31,13 +24,13 @@ func Login(arguments []string) error { provider := arguments[0] // Ask the server which oauth apps to log in against - providersRequest, err := apiRequest(http.MethodGet, "/login", nil) + providersRequest, err := apiRequest(session, http.MethodGet, "/login", nil) if err != nil { return err } - providersResponse, err := apiClient.Do(providersRequest) + providersResponse, err := session.Client.Do(providersRequest) if err != nil { return fmt.Errorf("the server could not be reached: %w", err) @@ -46,9 +39,7 @@ func Login(arguments []string) error { defer providersResponse.Body.Close() if providersResponse.StatusCode != http.StatusOK { - message, _ := io.ReadAll(io.LimitReader(providersResponse.Body, 4096)) - - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return serverError(providersResponse) } providers := struct { @@ -68,14 +59,14 @@ func Login(arguments []string) error { switch provider { case "github": clientId = providers.GithubClientId - deviceCodeUrl = githubBase + "/login/device/code" - pollUrl = githubBase + "/login/oauth/access_token" + deviceCodeUrl = session.GithubBase + "/login/device/code" + pollUrl = session.GithubBase + "/login/oauth/access_token" scope = "user:email" case "gitlab": clientId = providers.GitlabClientId - deviceCodeUrl = gitlabBase + "/oauth/authorize_device" - pollUrl = gitlabBase + "/oauth/token" + deviceCodeUrl = session.GitlabBase + "/oauth/authorize_device" + pollUrl = session.GitlabBase + "/oauth/token" scope = "read_user" } @@ -133,34 +124,33 @@ func Login(arguments []string) error { enterAt = code.VerificationUriComplete } - fmt.Printf("Copy your one-time code: %s\n", code.UserCode) - fmt.Printf("Then enter it at %s\n", enterAt) - fmt.Println("Press enter to open the browser.") - - // The read sits in a goroutine so an unpressed key never stalls the - // poll: the code may just as well be entered on another device. The - // stream is captured here, because the goroutine outlives the read and - // must not touch os.Stdin once something else may have replaced it. - prompt := os.Stdin + fmt.Fprintf(session.Out, "Copy your one-time code: %s\n", code.UserCode) + fmt.Fprintf(session.Out, "Then enter it at %s\n", enterAt) + fmt.Fprintln(session.Out, "Press enter to open the browser.") go func() { - _, err := bufio.NewReader(prompt).ReadString('\n') + _, err := bufio.NewReader(session.In).ReadString('\n') if err == nil { - openBrowser(enterAt) + session.OpenBrowser(enterAt) } }() // Poll until the code is entered deadline := time.Now().Add(time.Duration(code.ExpiresIn) * time.Second) + const defaultPollInterval = 5 interval := code.Interval + if interval <= 0 { + interval = defaultPollInterval + } + accessToken := "" for accessToken == "" { - time.Sleep(time.Duration(max(interval, minimumPollInterval)) * time.Second) + time.Sleep(time.Duration(interval) * time.Second) if time.Now().After(deadline) { return errors.New("the code expired before it was entered, run login again") @@ -191,7 +181,6 @@ func Login(arguments []string) error { poll := struct { AccessToken string `json:"access_token"` Error string `json:"error"` - Interval int `json:"interval"` }{} err = json.NewDecoder(pollResponse.Body).Decode(&poll) @@ -213,11 +202,7 @@ func Login(arguments []string) error { case "authorization_pending": case "slow_down": - if poll.Interval > 0 { - interval = poll.Interval - } else { - interval += 5 - } + interval += 5 case "expired_token": return errors.New("the code expired before it was entered, run login again") @@ -240,7 +225,7 @@ func Login(arguments []string) error { return err } - loginRequest, err := apiRequest(http.MethodPost, "/login", bytes.NewReader(loginBody)) + loginRequest, err := apiRequest(session, http.MethodPost, "/login", bytes.NewReader(loginBody)) if err != nil { return err @@ -248,7 +233,7 @@ func Login(arguments []string) error { loginRequest.Header.Set("Content-Type", "application/json") - loginResponse, err := apiClient.Do(loginRequest) + loginResponse, err := session.Client.Do(loginRequest) if err != nil { return fmt.Errorf("the server could not be reached: %w", err) @@ -257,9 +242,7 @@ func Login(arguments []string) error { defer loginResponse.Body.Close() if loginResponse.StatusCode != http.StatusOK { - message, _ := io.ReadAll(io.LimitReader(loginResponse.Body, 4096)) - - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return serverError(loginResponse) } login := struct { @@ -296,7 +279,7 @@ func Login(arguments []string) error { return err } - fmt.Printf("Logged in as %s.\n", login.Email) + fmt.Fprintf(session.Out, "Logged in as %s.\n", login.Email) return nil } diff --git a/internal/commands/login_test.go b/internal/commands/login_test.go index 40ea715..e51e936 100644 --- a/internal/commands/login_test.go +++ b/internal/commands/login_test.go @@ -1,6 +1,7 @@ package commands import ( + "bytes" "encoding/json" "fmt" "net/http" @@ -11,7 +12,7 @@ import ( "time" ) -func fakeProviderForLogin(t *testing.T, provider string, deviceInterval int, pollAnswers []string) *[]time.Time { +func fakeProviderForLogin(t *testing.T, provider string, deviceInterval int, pollAnswers []string) (*[]time.Time, string) { t.Helper() devicePath := "/login/device/code" @@ -100,24 +101,10 @@ func fakeProviderForLogin(t *testing.T, provider string, deviceInterval int, pol t.Cleanup(server.Close) - if provider == "gitlab" { - previousBase := gitlabBase - - gitlabBase = server.URL - - t.Cleanup(func() { gitlabBase = previousBase }) - } else { - previousBase := githubBase - - githubBase = server.URL - - t.Cleanup(func() { githubBase = previousBase }) - } - - return polledAt + return polledAt, server.URL } -func fakeSuperstack(t *testing.T) { +func fakeSuperstack(t *testing.T) (Session, *bytes.Buffer) { t.Helper() mux := http.NewServeMux() @@ -152,9 +139,11 @@ func fakeSuperstack(t *testing.T) { t.Cleanup(server.Close) - chosenApiBase = server.URL + out := &bytes.Buffer{} + session := NewSession(server.URL, "test", strings.NewReader(""), out) + session.OpenBrowser = func(url string) {} - t.Cleanup(func() { chosenApiBase = "" }) + return session, out } func TestLogin(t *testing.T) { @@ -200,7 +189,7 @@ func TestLogin(t *testing.T) { `{"error": "slow_down", "interval": 1}`, `{"access_token": "gho_test"}`, }, - wantPollGap: time.Second, + wantPollGap: 6 * time.Second, }, { name: "gitlab slowed down without an interval", @@ -209,7 +198,7 @@ func TestLogin(t *testing.T) { `{"error": "slow_down"}`, `{"access_token": "glpat-test"}`, }, - wantPollGap: 5 * time.Second, + wantPollGap: 6 * time.Second, }, { name: "the device interval is honored", @@ -251,21 +240,22 @@ func TestLogin(t *testing.T) { t.Run(test.name, func(t *testing.T) { isolateKeyStorage(t) - answerOnStdin(t, "") - - captureBrowserOpens(t) - - previousMinimum := minimumPollInterval - - minimumPollInterval = 0 + deviceInterval := test.deviceInterval - t.Cleanup(func() { minimumPollInterval = previousMinimum }) + if deviceInterval == 0 { + deviceInterval = 1 + } - polledAt := fakeProviderForLogin(t, test.provider, test.deviceInterval, test.pollAnswers) + polledAt, providerBase := fakeProviderForLogin(t, test.provider, deviceInterval, test.pollAnswers) + session, _ := fakeSuperstack(t) - fakeSuperstack(t) + if test.provider == "gitlab" { + session.GitlabBase = providerBase + } else { + session.GithubBase = providerBase + } - err := Login([]string{test.provider}) + err := Login(session, []string{test.provider}) if test.wantPollGap > 0 { if len(*polledAt) < 2 { @@ -329,23 +319,14 @@ func TestLogin(t *testing.T) { func TestLoginOpensTheBrowserOnEnter(t *testing.T) { isolateKeyStorage(t) - previousMinimum := minimumPollInterval - - minimumPollInterval = 0 - - t.Cleanup(func() { minimumPollInterval = previousMinimum }) - - fakeProviderForLogin(t, "gitlab", 0, []string{`{"access_token": "glpat-test"}`}) + _, providerBase := fakeProviderForLogin(t, "gitlab", 1, []string{`{"access_token": "glpat-test"}`}) + session, _ := fakeSuperstack(t) + session.GitlabBase = providerBase + session.In = strings.NewReader("\n") + browserOpens := make(chan string, 1) + session.OpenBrowser = func(url string) { browserOpens <- url } - fakeSuperstack(t) - - browserOpens := captureBrowserOpens(t) - - answerOnStdin(t, "\n") - - _, err := captureStdout(t, func() error { - return Login([]string{"gitlab"}) - }) + err := Login(session, []string{"gitlab"}) if err != nil { t.Fatal(err) @@ -374,7 +355,7 @@ func TestLoginRequiresAProvider(t *testing.T) { } for _, test := range tests { - err := Login(test.arguments) + err := Login(Session{}, test.arguments) if err == nil || !strings.Contains(err.Error(), "a provider: github or gitlab") { t.Errorf("%s: error = %v, want the provider hint", test.name, err) @@ -395,11 +376,10 @@ func TestLoginProviderNotOffered(t *testing.T) { t.Cleanup(server.Close) - chosenApiBase = server.URL - - t.Cleanup(func() { chosenApiBase = "" }) + out := &bytes.Buffer{} + session := NewSession(server.URL, "test", strings.NewReader(""), out) - err := Login([]string{"gitlab"}) + err := Login(session, []string{"gitlab"}) if err == nil || !strings.Contains(err.Error(), "offers no gitlab login") { t.Fatalf("error = %v, want it to say the server offers no gitlab login", err) diff --git a/internal/commands/logout.go b/internal/commands/logout.go index 93dd681..4a4ad5d 100644 --- a/internal/commands/logout.go +++ b/internal/commands/logout.go @@ -3,14 +3,13 @@ package commands import ( "errors" "fmt" - "io" "io/fs" "net/http" "os" "strings" ) -func Logout(arguments []string) error { +func Logout(session Session, arguments []string) error { if len(arguments) != 0 { return errors.New("logout takes no arguments") @@ -25,7 +24,7 @@ func Logout(arguments []string) error { keyBytes, err := os.ReadFile(path) if errors.Is(err, fs.ErrNotExist) { - fmt.Println("Not logged in.") + fmt.Fprintln(session.Out, "Not logged in.") return nil } @@ -35,7 +34,7 @@ func Logout(arguments []string) error { // Revoke on the server first, keeping the key on any failure so another // logout can retry; a forgotten key can never be revoked - revokeRequest, err := apiRequest(http.MethodPost, "/logout", nil) + revokeRequest, err := apiRequest(session, http.MethodPost, "/logout", nil) if err != nil { return err @@ -43,7 +42,7 @@ func Logout(arguments []string) error { revokeRequest.Header.Set("Authorization", "Bearer "+strings.TrimSpace(string(keyBytes))) - revokeResponse, err := apiClient.Do(revokeRequest) + revokeResponse, err := session.Client.Do(revokeRequest) if err != nil { return fmt.Errorf("you are still logged in: the server could not be reached: %w", err) @@ -52,15 +51,7 @@ func Logout(arguments []string) error { defer revokeResponse.Body.Close() if revokeResponse.StatusCode != http.StatusNoContent { - message, _ := io.ReadAll(io.LimitReader(revokeResponse.Body, 4096)) - - detail := strings.TrimSpace(string(message)) - - if detail == "" { - detail = revokeResponse.Status - } - - return fmt.Errorf("you are still logged in: the server said: %s", detail) + return fmt.Errorf("you are still logged in: %w", serverError(revokeResponse)) } err = os.Remove(path) @@ -69,7 +60,7 @@ func Logout(arguments []string) error { return err } - fmt.Println("Logged out.") + fmt.Fprintln(session.Out, "Logged out.") return nil } diff --git a/internal/commands/logout_test.go b/internal/commands/logout_test.go index 12f6205..cf636cc 100644 --- a/internal/commands/logout_test.go +++ b/internal/commands/logout_test.go @@ -1,6 +1,7 @@ package commands import ( + "bytes" "net/http" "net/http/httptest" "os" @@ -18,15 +19,18 @@ func TestLogout(t *testing.T) { wantError string wantRevocation bool wantKeyKept bool + wantShown string }{ { name: "revokes and forgets the stored key", storedKey: "ssk_test", revokeStatus: http.StatusNoContent, wantRevocation: true, + wantShown: "Logged out.\n", }, { - name: "nothing stored", + name: "nothing stored", + wantShown: "Not logged in.\n", }, { name: "server refuses the revocation", @@ -67,9 +71,8 @@ func TestLogout(t *testing.T) { server.Close() } - chosenApiBase = server.URL - - t.Cleanup(func() { chosenApiBase = "" }) + out := &bytes.Buffer{} + session := NewSession(server.URL, "test", strings.NewReader(""), out) path, err := keyPath() @@ -91,7 +94,7 @@ func TestLogout(t *testing.T) { } } - err = Logout(nil) + err = Logout(session, nil) if test.wantError != "" { if err == nil || !strings.Contains(err.Error(), test.wantError) { @@ -118,6 +121,10 @@ func TestLogout(t *testing.T) { if !test.wantKeyKept && !os.IsNotExist(statError) { t.Error("the stored key still exists after logout") } + + if out.String() != test.wantShown { + t.Errorf("output = %q, want %q", out.String(), test.wantShown) + } }) } } diff --git a/internal/commands/member_add.go b/internal/commands/member_add.go index c21b2e9..c7a4a6b 100644 --- a/internal/commands/member_add.go +++ b/internal/commands/member_add.go @@ -5,13 +5,11 @@ import ( "encoding/json" "errors" "fmt" - "io" "net/http" "strconv" - "strings" ) -func MemberAdd(arguments []string) error { +func MemberAdd(session Session, arguments []string) error { if len(arguments) != 2 || arguments[0] == "" { return errors.New("member add takes an email address and a fleet id") @@ -31,7 +29,7 @@ func MemberAdd(arguments []string) error { return err } - request, err := authenticatedRequest(http.MethodPost, + request, err := authenticatedRequest(session, http.MethodPost, "/fleets/"+strconv.FormatInt(fleetId, 10)+"/members", bytes.NewReader(body)) if err != nil { @@ -40,7 +38,7 @@ func MemberAdd(arguments []string) error { request.Header.Set("Content-Type", "application/json") - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return fmt.Errorf("the server could not be reached: %w", err) @@ -49,12 +47,10 @@ func MemberAdd(arguments []string) error { defer response.Body.Close() if response.StatusCode != http.StatusNoContent { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return serverError(response) } - fmt.Printf("Gave %s access.\n", email) + fmt.Fprintf(session.Out, "Gave %s access.\n", email) return nil } diff --git a/internal/commands/member_add_test.go b/internal/commands/member_add_test.go index 91a8030..af5fa62 100644 --- a/internal/commands/member_add_test.go +++ b/internal/commands/member_add_test.go @@ -26,9 +26,9 @@ func TestMemberAdd(t *testing.T) { w.WriteHeader(http.StatusNoContent) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - err := MemberAdd([]string{"member@example.com", "3"}) + err := MemberAdd(session, []string{"member@example.com", "3"}) if err != nil { t.Fatal(err) @@ -38,6 +38,10 @@ func TestMemberAdd(t *testing.T) { t.Errorf("the server saw %q added at %q, want %q at %q", addedEmail, addedPath, "member@example.com", "/fleets/3/members") } + + if out.String() != "Gave member@example.com access.\n" { + t.Errorf("output = %q", out.String()) + } } func TestMemberAddArguments(t *testing.T) { @@ -53,7 +57,7 @@ func TestMemberAddArguments(t *testing.T) { } for _, test := range tests { - err := MemberAdd(test.arguments) + err := MemberAdd(Session{}, test.arguments) if err == nil || !strings.Contains(err.Error(), test.wantError) { t.Errorf("%s: error = %v, want it to mention %q", test.name, err, test.wantError) diff --git a/internal/commands/member_list.go b/internal/commands/member_list.go index bc79cdd..79c56f0 100644 --- a/internal/commands/member_list.go +++ b/internal/commands/member_list.go @@ -4,27 +4,13 @@ import ( "encoding/json" "errors" "fmt" - "io" "net/http" - "os" "strconv" - "strings" ) -func MemberList(arguments []string) error { +func MemberList(session Session, arguments []string) error { - jsonOutput := false - - positionals := []string{} - - for _, argument := range arguments { - if argument == "--json" { - jsonOutput = true - continue - } - - positionals = append(positionals, argument) - } + positionals, jsonOutput := takeJsonFlag(arguments) if len(positionals) != 1 { return errors.New("member list takes a fleet id") @@ -36,14 +22,14 @@ func MemberList(arguments []string) error { return errors.New("the fleet id is the number shown by fleet list") } - request, err := authenticatedRequest(http.MethodGet, + request, err := authenticatedRequest(session, http.MethodGet, "/fleets/"+strconv.FormatInt(fleetId, 10)+"/members", nil) if err != nil { return err } - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return fmt.Errorf("the server could not be reached: %w", err) @@ -51,14 +37,8 @@ func MemberList(arguments []string) error { defer response.Body.Close() - body, err := io.ReadAll(response.Body) - - if err != nil { - return err - } - if response.StatusCode != http.StatusOK { - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(body))) + return serverError(response) } people := struct { @@ -66,14 +46,14 @@ func MemberList(arguments []string) error { Members []string `json:"members"` }{} - err = json.Unmarshal(body, &people) + err = json.NewDecoder(response.Body).Decode(&people) if err != nil { return err } if jsonOutput { - return json.NewEncoder(os.Stdout).Encode(people) + return json.NewEncoder(session.Out).Encode(people) } emailWidth := max(len("EMAIL"), len(people.Owner)) @@ -82,12 +62,12 @@ func MemberList(arguments []string) error { emailWidth = max(emailWidth, len(email)) } - fmt.Printf("%-*s %s\n", emailWidth, "EMAIL", "ROLE") + fmt.Fprintf(session.Out, "%-*s %s\n", emailWidth, "EMAIL", "ROLE") - fmt.Printf("%-*s owner\n", emailWidth, people.Owner) + fmt.Fprintf(session.Out, "%-*s owner\n", emailWidth, people.Owner) for _, email := range people.Members { - fmt.Printf("%-*s member\n", emailWidth, email) + fmt.Fprintf(session.Out, "%-*s member\n", emailWidth, email) } return nil diff --git a/internal/commands/member_list_test.go b/internal/commands/member_list_test.go index 9867d1b..7c15b0a 100644 --- a/internal/commands/member_list_test.go +++ b/internal/commands/member_list_test.go @@ -74,11 +74,11 @@ func TestMemberList(t *testing.T) { fmt.Fprint(w, test.people) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - printed, err := captureStdout(t, func() error { - return MemberList(test.arguments) - }) + err := MemberList(session, test.arguments) + + printed := out.String() if test.wantError != "" { if err == nil || !strings.Contains(err.Error(), test.wantError) { diff --git a/internal/commands/member_remove.go b/internal/commands/member_remove.go index 1b81a37..1167500 100644 --- a/internal/commands/member_remove.go +++ b/internal/commands/member_remove.go @@ -4,15 +4,13 @@ import ( "bufio" "errors" "fmt" - "io" "net/http" "net/url" - "os" "strconv" "strings" ) -func MemberRemove(arguments []string) error { +func MemberRemove(session Session, arguments []string) error { if len(arguments) != 2 || arguments[0] == "" { return errors.New("member remove takes an email address and a fleet id") @@ -26,7 +24,7 @@ func MemberRemove(arguments []string) error { return errors.New("the fleet id is the number shown by fleet list") } - fleets, err := fetchFleets() + fleets, err := fetchFleets(session) if err != nil { return err @@ -46,25 +44,25 @@ func MemberRemove(arguments []string) error { return errors.New("no such fleet") } - fmt.Printf("Take away %s's access to %q? [y/N] ", email, name) + fmt.Fprintf(session.Out, "Take away %s's access to %q? [y/N] ", email, name) - answer, _ := bufio.NewReader(os.Stdin).ReadString('\n') + answer, _ := bufio.NewReader(session.In).ReadString('\n') answer = strings.ToLower(strings.TrimSpace(answer)) if answer != "y" && answer != "yes" { - fmt.Println("Nothing changed.") + fmt.Fprintln(session.Out, "Nothing changed.") return nil } - request, err := authenticatedRequest(http.MethodDelete, + request, err := authenticatedRequest(session, http.MethodDelete, "/fleets/"+strconv.FormatInt(fleetId, 10)+"/members/"+url.PathEscape(email), nil) if err != nil { return err } - response, err := apiClient.Do(request) + response, err := session.Client.Do(request) if err != nil { return fmt.Errorf("the server could not be reached: %w", err) @@ -73,12 +71,10 @@ func MemberRemove(arguments []string) error { defer response.Body.Close() if response.StatusCode != http.StatusNoContent { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - - return fmt.Errorf("the server said: %s", strings.TrimSpace(string(message))) + return serverError(response) } - fmt.Printf("Removed access for %s.\n", email) + fmt.Fprintf(session.Out, "Removed access for %s.\n", email) return nil } diff --git a/internal/commands/member_remove_test.go b/internal/commands/member_remove_test.go index 1d4f34b..8a091b6 100644 --- a/internal/commands/member_remove_test.go +++ b/internal/commands/member_remove_test.go @@ -41,13 +41,13 @@ func TestMemberRemove(t *testing.T) { w.WriteHeader(http.StatusNoContent) }) - loggedInTestServer(t, mux) + session, out := loggedInSession(t, mux) - answerOnStdin(t, test.answer) + session.In = strings.NewReader(test.answer) - printed, err := captureStdout(t, func() error { - return MemberRemove([]string{test.email, "3"}) - }) + err := MemberRemove(session, []string{test.email, "3"}) + + printed := out.String() if err != nil { t.Fatal(err) @@ -82,7 +82,7 @@ func TestMemberRemoveArguments(t *testing.T) { } for _, test := range tests { - err := MemberRemove(test.arguments) + err := MemberRemove(Session{}, test.arguments) if err == nil || !strings.Contains(err.Error(), test.wantError) { t.Errorf("%s: error = %v, want it to mention %q", test.name, err, test.wantError) diff --git a/internal/commands/session.go b/internal/commands/session.go new file mode 100644 index 0000000..a50a8c3 --- /dev/null +++ b/internal/commands/session.go @@ -0,0 +1,35 @@ +package commands + +import ( + "io" + "net/http" + "time" +) + +const defaultApiBase = "https://supernext.siliconwitchery.com" +const defaultGithubBase = "https://github.com" +const defaultGitlabBase = "https://gitlab.com" + +type Session struct { + Base string + GithubBase string + GitlabBase string + Version string + Client *http.Client + In io.Reader + Out io.Writer + OpenBrowser func(url string) +} + +func NewSession(base string, version string, in io.Reader, out io.Writer) Session { + return Session{ + Base: base, + GithubBase: defaultGithubBase, + GitlabBase: defaultGitlabBase, + Version: version, + Client: &http.Client{Timeout: 30 * time.Second}, + In: in, + Out: out, + OpenBrowser: openBrowser, + } +} diff --git a/main.go b/main.go index 2080939..ec9cbdd 100644 --- a/main.go +++ b/main.go @@ -2,249 +2,92 @@ package main import ( "fmt" - "io" "os" - "strings" "github.com/siliconwitchery/superstack-cli/internal/commands" ) const version = "0.0.3" -type command struct { - name string - arguments string - summary string - run func(arguments []string) error -} - -type section struct { - title string - commands []command -} - -var sections = []section{ +var sections = []commands.Section{ { - title: "Getting started", - commands: []command{ - {name: "login", arguments: "", summary: "Log in with the selected provider", run: commands.Login}, - {name: "logout", summary: "Log out of your account", run: commands.Logout}, + Title: "Getting started", + Commands: []commands.Command{ + {Name: "login", Arguments: "", Summary: "Log in with the selected provider", Run: commands.Login}, + {Name: "logout", Summary: "Log out of your account", Run: commands.Logout}, }, }, { - title: "Fleets", - commands: []command{ - {name: "fleet create", arguments: "", summary: "Create a fleet", run: commands.FleetCreate}, - {name: "fleet list", arguments: "[--json]", summary: "List the fleets you can reach", run: commands.FleetList}, - {name: "fleet rename", arguments: " ", summary: "Rename a fleet", run: commands.FleetRename}, - {name: "fleet transfer", arguments: " ", summary: "Hand a fleet to a new owner", run: commands.FleetTransfer}, - {name: "fleet delete", arguments: "", summary: "Delete a fleet and factory reset its devices", run: commands.FleetDelete}, + Title: "Fleets", + Commands: []commands.Command{ + {Name: "fleet create", Arguments: "", Summary: "Create a fleet", Run: commands.FleetCreate}, + {Name: "fleet list", Arguments: "[--json]", Summary: "List the fleets you can reach", Run: commands.FleetList}, + {Name: "fleet rename", Arguments: " ", Summary: "Rename a fleet", Run: commands.FleetRename}, + {Name: "fleet transfer", Arguments: " ", Summary: "Hand a fleet to a new owner", Run: commands.FleetTransfer}, + {Name: "fleet delete", Arguments: "", Summary: "Delete a fleet and factory reset its devices", Run: commands.FleetDelete}, }, }, { - title: "Devices", - commands: []command{ - {name: "device claim", arguments: " [name]", summary: "Claim a device into a fleet, then press its button", run: commands.DeviceClaim}, - {name: "device list", arguments: "[fleet_id] [--json]", summary: "List devices, their state, and when they were last seen", run: commands.DeviceList}, - {name: "device rename", arguments: " ", summary: "Rename a device", run: commands.DeviceRename}, - {name: "device release", arguments: "", summary: "Unpair a device from its fleet and factory reset it", run: commands.DeviceRelease}, - {name: "device start", arguments: "", summary: "Run the code on the target"}, - {name: "device stop", arguments: "", summary: "Halt the code on the target"}, - {name: "device restart", arguments: "", summary: "Restart the code on the target"}, + Title: "Devices", + Commands: []commands.Command{ + {Name: "device claim", Arguments: " [name]", Summary: "Claim a device into a fleet, then press its button", Run: commands.DeviceClaim}, + {Name: "device list", Arguments: "[fleet_id] [--json]", Summary: "List devices, their state, and when they were last seen", Run: commands.DeviceList}, + {Name: "device rename", Arguments: " ", Summary: "Rename a device", Run: commands.DeviceRename}, + {Name: "device release", Arguments: "", Summary: "Unpair a device from its fleet and factory reset it", Run: commands.DeviceRelease}, + {Name: "device start", Arguments: "", Summary: "Run the code on the target"}, + {Name: "device stop", Arguments: "", Summary: "Halt the code on the target"}, + {Name: "device restart", Arguments: "", Summary: "Restart the code on the target"}, }, }, { - title: "Files", - commands: []command{ - {name: "upload", arguments: " ...", summary: "Upload files or directories to the target"}, - {name: "download", arguments: " ", summary: "Download the target's files into "}, - {name: "dev", arguments: " ... [--log-file ]", summary: "Upload on every change, and tail"}, + Title: "Files", + Commands: []commands.Command{ + {Name: "upload", Arguments: " ...", Summary: "Upload files or directories to the target"}, + {Name: "download", Arguments: " ", Summary: "Download the target's files into "}, + {Name: "dev", Arguments: " ... [--log-file ]", Summary: "Upload on every change, and tail"}, }, }, { - title: "Logs", - commands: []command{ - {name: "tail", arguments: " [-n num] [--log-file ]", summary: "Stream the target's log as it arrives"}, + Title: "Logs", + Commands: []commands.Command{ + {Name: "tail", Arguments: " [-n num] [--log-file ]", Summary: "Stream the target's log as it arrives"}, }, }, { - title: "People", - commands: []command{ - {name: "member add", arguments: " ", summary: "Give someone access to a fleet", run: commands.MemberAdd}, - {name: "member list", arguments: " [--json]", summary: "List the people who can reach a fleet", run: commands.MemberList}, - {name: "member remove", arguments: " ", summary: "Take away someone's access", run: commands.MemberRemove}, + Title: "People", + Commands: []commands.Command{ + {Name: "member add", Arguments: " ", Summary: "Give someone access to a fleet", Run: commands.MemberAdd}, + {Name: "member list", Arguments: " [--json]", Summary: "List the people who can reach a fleet", Run: commands.MemberList}, + {Name: "member remove", Arguments: " ", Summary: "Take away someone's access", Run: commands.MemberRemove}, }, }, { - title: "Keys", - commands: []command{ - {name: "key create", arguments: "