Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
164 changes: 143 additions & 21 deletions cmd/login.go
Original file line number Diff line number Diff line change
@@ -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"
Expand All @@ -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 (
Expand Down Expand Up @@ -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() {
Expand All @@ -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.")
}
}

Expand Down
147 changes: 147 additions & 0 deletions cmd/login_test.go
Original file line number Diff line number Diff line change
@@ -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))
}
}
Loading