From 37399b439f060e1661fe31b09d49cdd9f2340d41 Mon Sep 17 00:00:00 2001 From: Caio Pizzol <33255434+caiopizzol@users.noreply.github.com> Date: Wed, 19 Aug 2026 18:41:48 -0300 Subject: [PATCH] feat: add headless login --- cmd/login.go | 164 ++++++++++++++++++++++++++++++++++++++++------ cmd/login_test.go | 147 +++++++++++++++++++++++++++++++++++++++++ 2 files changed, 290 insertions(+), 21 deletions(-) create mode 100644 cmd/login_test.go diff --git a/cmd/login.go b/cmd/login.go index f71d34b..5e1cea6 100644 --- a/cmd/login.go +++ b/cmd/login.go @@ -1,14 +1,19 @@ package cmd import ( + "bufio" "context" "encoding/base64" "encoding/json" + "errors" "fmt" + "io" "net/http" + "net/url" "os" "os/exec" "runtime" + "strings" "time" "github.com/amp-labs/cli/clerk" @@ -23,6 +28,11 @@ const ( ServerPort = 3535 ) +var ( + errInvalidLoginCallback = errors.New("invalid callback URL") + errMissingLoginPayload = errors.New("callback URL is missing its login payload") +) + type handler struct{} const ( @@ -103,15 +113,137 @@ func processLogin(ctx context.Context, payload []byte, write bool) (string, stri const ReadHeaderTimeoutSeconds = 3 -// loginCmd represents the login command. -var loginCmd = &cobra.Command{ //nolint:gochecknoglobals - Use: "login", - Short: "Log into an Ampersand account", - Long: "Log into an Ampersand account.", - Run: func(cmd *cobra.Command, args []string) { - DoLogout(false) - doLogin() - }, +func newLoginCmd( + logout func(bool), + login func(), + headlessLogin func(context.Context), +) *cobra.Command { + var headless bool + + cmd := &cobra.Command{ + Use: "login", + Short: "Log into an Ampersand account", + Long: "Log into an Ampersand account.", + Run: func(cmd *cobra.Command, args []string) { + logout(false) + + if headless { + headlessLogin(cmd.Context()) + + return + } + + login() + }, + } + + cmd.Flags().BoolVar(&headless, "headless", false, "Log in using a browser on another machine") + + return cmd +} + +var loginCmd = newLoginCmd(DoLogout, doLogin, doHeadlessLogin) //nolint:gochecknoglobals + +func doHeadlessLogin(ctx context.Context) { + // Reuse the hosted page's existing callback so this flow needs no new auth endpoint. + logger.Infof("Open %s in a browser.", getLoginURL()) + logger.Info("After signing in, copy the localhost URL from your browser and paste it here.") + fmt.Fprint(os.Stdout, "Paste callback URL: ") + + callback, err := readHiddenInput(os.Stdin) + + fmt.Fprintln(os.Stdout) + + if err != nil { + logger.FatalErr("Unable to read callback URL", err) + } + + payload, err := parseLoginCallback(string(callback)) + if err != nil { + logger.FatalErr("Unable to complete login", err) + } + + _, email, err := processLogin(ctx, payload, true) + if err != nil { + logger.FatalErr("Unable to complete login", err) + } + + logger.Info("Successfully logged in as " + email) +} + +func readHiddenInput(input *os.File) ([]byte, error) { + state, err := term.MakeRaw(int(input.Fd())) + if err != nil { + return nil, err + } + + defer func() { + _ = term.Restore(int(input.Fd()), state) + }() + + return readInputLine(input) +} + +func readInputLine(reader io.Reader) ([]byte, error) { + const deleteCharacter = '\x7f' + + var result []byte + + buffered := bufio.NewReader(reader) + + for { + character, err := buffered.ReadByte() + if err == nil { + switch character { + case '\r', '\n': + return result, nil + case '\b', deleteCharacter: + if len(result) > 0 { + result = result[:len(result)-1] + } + default: + result = append(result, character) + } + + continue + } + + if errors.Is(err, io.EOF) && len(result) > 0 { + return result, nil + } + + return result, err + } +} + +func parseLoginCallback(callback string) ([]byte, error) { + parsed, err := url.Parse(strings.TrimSpace(callback)) + if err != nil { + return nil, fmt.Errorf("invalid callback URL: %w", err) + } + + if parsed.Scheme != "http" || + parsed.Host != fmt.Sprintf("localhost:%d", ServerPort) || + parsed.Path != "/done" || parsed.Fragment != "" { + return nil, errInvalidLoginCallback + } + + encoded, found := strings.CutPrefix(parsed.RawQuery, "p=") + if !found || encoded == "" || strings.Contains(encoded, "&") { + return nil, errMissingLoginPayload + } + + encoded, err = url.PathUnescape(encoded) + if err != nil { + return nil, fmt.Errorf("invalid callback payload: %w", err) + } + + payload, err := base64.StdEncoding.DecodeString(encoded) + if err != nil { + return nil, fmt.Errorf("invalid callback payload: %w", err) + } + + return payload, nil } func doLogin() { @@ -125,18 +257,8 @@ func doLogin() { if hasBrowser { openBrowser(fmt.Sprintf("http://localhost:%d", ServerPort)) } else { - link := getLoginURL() - - linkMsg := fmt.Sprintf("No browser detected, please open %s in your browser to log in.", link) - localhostMsg := fmt.Sprintf("NOTE: the login page will redirect to http://localhost:%d/...", ServerPort) - - logger.Info(linkMsg) - logger.Info() - logger.Info(localhostMsg) - logger.Info("If this URL isn't accessible (e.g. you're using a remote server),") - logger.Info("the credentials won't be saved. It's best to run this command") - logger.Info("on a machine with a browser, but you can also overcome this using") - logger.Info("SSH port forwarding or a proxy.") + logger.Info("No browser detected.") + logger.Info("Stop this command and run `amp login --headless` to sign in from another machine.") } } diff --git a/cmd/login_test.go b/cmd/login_test.go new file mode 100644 index 0000000..569e916 --- /dev/null +++ b/cmd/login_test.go @@ -0,0 +1,147 @@ +package cmd + +import ( + "context" + "encoding/base64" + "fmt" + "strings" + "testing" +) + +func TestLoginCommandUsesHeadlessRunner(t *testing.T) { + t.Parallel() + + loggedOut := false + openedBrowser := false + ranHeadless := false + cmd := newLoginCmd( + func(showLogs bool) { + if showLogs { + t.Fatal("logout logs were enabled") + } + + loggedOut = true + }, + func() { openedBrowser = true }, + func(context.Context) { ranHeadless = true }, + ) + cmd.SetArgs([]string{"--headless"}) + + err := cmd.Execute() + if err != nil { + t.Fatalf("execute login command: %v", err) + } + + if !loggedOut { + t.Fatal("existing login was not cleared") + } + + if openedBrowser { + t.Fatal("browser login ran in headless mode") + } + + if !ranHeadless { + t.Fatal("headless login did not run") + } +} + +func TestLoginCommandUsesBrowserRunnerByDefault(t *testing.T) { + t.Parallel() + + openedBrowser := false + ranHeadless := false + cmd := newLoginCmd( + func(bool) {}, + func() { openedBrowser = true }, + func(context.Context) { ranHeadless = true }, + ) + + err := cmd.Execute() + if err != nil { + t.Fatalf("execute login command: %v", err) + } + + if !openedBrowser { + t.Fatal("browser login did not run") + } + + if ranHeadless { + t.Fatal("headless login ran without --headless") + } +} + +func TestParseLoginCallback(t *testing.T) { + t.Parallel() + + payload := []byte(`{"cookies":{"session":"test"}}`) + callback := fmt.Sprintf( + "http://localhost:%d/done?p=%s", + ServerPort, + base64.StdEncoding.EncodeToString(payload), + ) + + got, err := parseLoginCallback(" " + callback + "\n") + if err != nil { + t.Fatalf("parse callback: %v", err) + } + + if string(got) != string(payload) { + t.Fatalf("payload = %q, want %q", got, payload) + } +} + +func TestParseLoginCallbackPreservesBase64Plus(t *testing.T) { + t.Parallel() + + got, err := parseLoginCallback(fmt.Sprintf("http://localhost:%d/done?p=+w==", ServerPort)) + if err != nil { + t.Fatalf("parse callback: %v", err) + } + + if len(got) != 1 || got[0] != 0xfb { + t.Fatalf("payload = %v, want [251]", got) + } +} + +func TestParseLoginCallbackRejectsInvalidInput(t *testing.T) { + t.Parallel() + + testCases := map[string]string{ + "wrong scheme": fmt.Sprintf("https://localhost:%d/done?p=e30=", ServerPort), + "wrong host": "http://example.com/done?p=e30=", + "wrong port": "http://localhost:9999/done?p=e30=", + "wrong path": fmt.Sprintf("http://localhost:%d/other?p=e30=", ServerPort), + "missing value": fmt.Sprintf("http://localhost:%d/done", ServerPort), + "extra query": fmt.Sprintf("http://localhost:%d/done?p=e30=&other=value", ServerPort), + "invalid base64": fmt.Sprintf( + "http://localhost:%d/done?p=not-base64", + ServerPort, + ), + } + + for name, callback := range testCases { + t.Run(name, func(t *testing.T) { + t.Parallel() + + _, err := parseLoginCallback(callback) + if err == nil { + t.Fatal("expected callback to be rejected") + } + }) + } +} + +func TestReadInputLineAcceptsLongCallback(t *testing.T) { + t.Parallel() + + want := strings.Repeat("a", 32_000) + + got, err := readInputLine(strings.NewReader(want + "\n")) + if err != nil { + t.Fatalf("read input: %v", err) + } + + if string(got) != want { + t.Fatalf("input length = %d, want %d", len(got), len(want)) + } +}