diff --git a/CLAUDE.md b/CLAUDE.md index 877ee04f0..21de059f7 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -53,7 +53,7 @@ go test -v -run TestName ./path/to/package ### Hybrid CLI System The CLI operates as a wrapper around a legacy PHP CLI: -- Go layer: Handles new commands (init, list, version, config:install, project:convert) and core infrastructure +- Go layer: Handles new commands (init, list, version, config:install, project:convert, the auth commands) and core infrastructure, including authentication - PHP layer: Legacy commands are proxied through `internal/legacy/CLIWrapper` - The PHP CLI (platform.phar) is embedded at build time via go:embed - An index of legacy commands (commands.json, from `list --all --format=json`) is embedded too, so the Go layer can resolve abbreviations like `p:init` in the same way as Symfony Console @@ -67,7 +67,7 @@ The CLI operates as a wrapper around a legacy PHP CLI: **Commands**: `commands/` - `root.go`: Root command that sets up the Cobra CLI and delegates to legacy CLI when needed -- Native Go commands: init, list, version, config:install, project:convert, completion +- Native Go commands: init, list, version, config:install, project:convert, completion, and auth:browser-login (login), auth:api-token-login, auth:logout (logout), auth:token - Unrecognized commands are passed to the legacy PHP CLI **Configuration**: `internal/config/` @@ -90,8 +90,10 @@ The CLI operates as a wrapper around a legacy PHP CLI: - Handles authentication, organizations, and resource management **Authentication**: `internal/auth/` -- JWT handling and OAuth2 flow -- Custom transport for API authentication +- Go is the only component that stores or refreshes credentials. `auth.Manager` resolves tokens (API tokens, `api.access_token`, stored sessions) and refreshes them under a per-session flock (`/auth/.lock`), re-reading the store under the lock because refresh tokens rotate +- `internal/auth/store`: one entry per session ID, in the system keychain (go-keyring) or in `/auth/.json` +- Auth settings are read by `config.Auth()` with the legacy CLI's precedence: embedded config, the user's `config.yaml`, env vars +- The legacy CLI gets tokens and auth state by running the hidden `auth:internal token|status` command (via `WRAPPER_EXECUTABLE`), and Go runs the hidden PHP commands `auth:post-login` (SSH certificates and config) and `auth:export-sessions` (a one-time migration of the legacy storage, recorded in `auth/.migrated`) **Project Initialization**: `internal/init/` - AI-powered project configuration generation diff --git a/README.md b/README.md index d2eb9d1aa..c79dc0dfe 100644 --- a/README.md +++ b/README.md @@ -199,6 +199,9 @@ Environment variables include: This skips confirmation questions. - `UPSUN_CLI_SESSION_ID`: switch user session (default `default`). See also `upsun session:switch`. +- `UPSUN_CLI_API_DISABLE_CREDENTIAL_HELPERS=1`: store credentials in files under + `~/.upsun-cli/auth/` instead of the system keychain. Files are also used when + no keychain is available. - `UPSUN_CLI_AUTO_LOAD_SSH_CERT=0`: disable automatically loading an SSH certificate when running login or SSH commands. - `UPSUN_CLI_SHELL_CONFIG_FILE`: the shell config file that `self:install` diff --git a/commands/auth.go b/commands/auth.go new file mode 100644 index 000000000..4e3044e9a --- /dev/null +++ b/commands/auth.go @@ -0,0 +1,419 @@ +package commands + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "os/exec" + "runtime" + "slices" + "strings" + "sync" + + "github.com/fatih/color" + "github.com/spf13/cobra" + "github.com/spf13/viper" + "golang.org/x/term" + + "github.com/upsun/cli/internal/auth" + "github.com/upsun/cli/internal/config" +) + +// authCommands returns the native auth commands that are enabled, which are listed alongside the legacy CLI's. +func authCommands(cnf *config.Config) []*cobra.Command { + all := []*cobra.Command{ + newAPITokenLoginCommand(cnf), + newAuthTokenCommand(cnf), + newBrowserLoginCommand(cnf), + newLogoutCommand(cnf), + } + return slices.DeleteFunc(all, func(c *cobra.Command) bool { + return slices.Contains(cnf.Application.DisabledCommands, c.Name()) || + slices.Contains(cnf.Application.WrappedDisabledCommands, c.Name()) + }) +} + +// addLegacyGlobalFlags accepts the legacy CLI's global options that the root command does not define. +func addLegacyGlobalFlags(c *cobra.Command) { + c.Flags().BoolP("no", "n", false, "Answer \"no\" to confirmation questions; disable interaction") + c.Flags().Bool("ansi", false, "Force ANSI output") + c.Flags().Bool("no-ansi", false, "Disable ANSI output") + for _, name := range []string{"no", "ansi", "no-ansi"} { + _ = c.Flags().MarkHidden(name) + } +} + +// applyLegacyGlobalFlags applies the options added by addLegacyGlobalFlags. +func applyLegacyGlobalFlags(c *cobra.Command) { + if no, _ := c.Flags().GetBool("no"); no { + viper.Set("no", true) + viper.Set("no-interaction", true) + } + if ansi, _ := c.Flags().GetBool("ansi"); ansi { + color.NoColor = false + } + if noANSI, _ := c.Flags().GetBool("no-ansi"); noANSI { + color.NoColor = true + } +} + +// newAuthManager creates the auth manager, which migrates the legacy CLI's sessions on first use. +func newAuthManager(cnf *config.Config, stderr io.Writer) (*auth.Manager, error) { + m, err := auth.NewManager(cnf, stderr) + if err != nil { + return nil, err + } + m.Migrator = &auth.Migrator{ + Export: func(ctx context.Context, del bool) ([]byte, error) { + var stdout, errOut bytes.Buffer + c := makeLegacyCLIWrapper(cnf, &stdout, &errOut, nil) + c.DisableInteraction = true + args := []string{"auth:export-sessions"} + if del { + args = append(args, "--delete") + } + if err := c.Exec(ctx, args...); err != nil { + return nil, fmt.Errorf("%w: %s", err, strings.TrimSpace(errOut.String())) + } + return stdout.Bytes(), nil + }, + DebugLog: debugLogf, + } + m.OnLoggedOut = func(id string) error { + return clearLegacySessionFiles(cnf, m.Settings.SessionID, []string{id}, false) + } + return m, nil +} + +// runLegacyAuthHook runs a hidden legacy CLI command that completes a login (SSH certificates and config). +func runLegacyAuthHook(cmd *cobra.Command, cnf *config.Config, args ...string) error { + c := makeLegacyCLIWrapper(cnf, cmd.OutOrStdout(), cmd.ErrOrStderr(), cmd.InOrStdin()) + return c.Exec(cmd.Context(), args...) +} + +// isInteractive reports whether questions may be asked, in the same way as the legacy CLI (Symfony Console). +func isInteractive(cmd *cobra.Command) bool { + if viper.GetBool("no-interaction") { + return false + } + if _, ok := os.LookupEnv("SHELL_INTERACTIVE"); ok { + return true + } + f, ok := cmd.InOrStdin().(*os.File) + return ok && term.IsTerminal(int(f.Fd())) +} + +// stdinReaders shares one buffered reader per input, so that consecutive questions do not lose buffered input. +var ( + stdinReaders = map[io.Reader]*bufio.Reader{} + stdinReadersMu sync.Mutex +) + +func readLine(r io.Reader) (string, error) { + stdinReadersMu.Lock() + br, ok := stdinReaders[r] + if !ok { + br = bufio.NewReader(r) + stdinReaders[r] = br + } + stdinReadersMu.Unlock() + line, err := br.ReadString('\n') + if err != nil && (!errors.Is(err, io.EOF) || line == "") { + return "", err + } + return strings.TrimRight(line, "\r\n"), nil +} + +// confirm asks a yes/no question. +func confirm(cmd *cobra.Command, question string, def bool) (bool, error) { + if viper.GetBool("yes") { + return true, nil + } + if viper.GetBool("no") { + return false, nil + } + if !isInteractive(cmd) { + return def, nil + } + hint := "Y/n" + if !def { + hint = "y/N" + } + for { + fmt.Fprintf(cmd.ErrOrStderr(), "%s [%s] ", question, hint) + answer, err := readLine(cmd.InOrStdin()) + if err != nil { + return false, err + } + switch strings.ToLower(strings.TrimSpace(answer)) { + case "": + return def, nil + case "y", "yes": + return true, nil + case "n", "no": + return false, nil + } + } +} + +// readSecret reads a line without echoing it, if the input is a terminal. +func readSecret(cmd *cobra.Command) (string, error) { + if f, ok := cmd.InOrStdin().(*os.File); ok && term.IsTerminal(int(f.Fd())) { + b, err := term.ReadPassword(int(f.Fd())) + fmt.Fprintln(cmd.ErrOrStderr()) + return string(b), err + } + return readLine(cmd.InOrStdin()) +} + +// exitError ends a command with an exit code, after its message has been printed. +type exitError struct{ code int } + +func (e *exitError) Error() string { return fmt.Sprintf("exit code %d", e.code) } + +// nonInteractiveAuthHelp matches the legacy CLI's Login::getNonInteractiveAuthHelp(). +func nonInteractiveAuthHelp(cnf *config.Config) string { + return fmt.Sprintf("To authenticate non-interactively, configure an API token using the %s environment variable.", + color.YellowString(cnf.Application.EnvPrefix+"TOKEN")) +} + +// sessionAdvice returns lines about the current session, if there are several, as in the legacy CLI. +func sessionAdvice(ctx context.Context, cnf *config.Config, m *auth.Manager, changeVerb string) []string { + ids, err := m.SessionIDs(ctx) + if err != nil { + debugLogf("Failed to list sessions: %s", err) + } + if m.Settings.SessionID == "default" && len(ids) <= 1 { + return nil + } + lines := []string{fmt.Sprintf("The current session ID is: %s", color.GreenString(m.Settings.SessionID))} + if !m.Settings.SessionIDFromEnv { + lines = append(lines, fmt.Sprintf("%s: %s", changeVerb, + color.GreenString(cnf.Application.Executable+" session:switch"))) + } + return lines +} + +// handleLoginRequired offers a browser login if possible, as the legacy CLI's AutoLoginListener did. +// It returns nil if the user logged in, or an *exitError (code 3) after printing why login is required. +func handleLoginRequired(cmd *cobra.Command, cnf *config.Config, m *auth.Manager, lerr *auth.LoginRequiredError) error { + stderr := cmd.ErrOrStderr() + if lerr.Notice != "" { + fmt.Fprintln(stderr, color.YellowString(lerr.Notice)) + fmt.Fprintln(stderr) + } + if isInteractive(cmd) && canOpenURLs("") { + fmt.Fprintln(stderr, lerr.Message()) + fmt.Fprintln(stderr) + if advice := sessionAdvice(cmd.Context(), cnf, m, "To switch sessions, run"); advice != nil { + fmt.Fprintln(stderr, strings.Join(advice, "\n")) + fmt.Fprintln(stderr) + } + ok, err := confirm(cmd, "Log in via a browser?", true) + if err != nil { + return err + } + if ok { + fmt.Fprintln(stderr) + opts := &browserLoginOptions{methods: lerr.AuthMethods, maxAge: lerr.MaxAge} + err := runBrowserLogin(cmd, cnf, m, opts) + fmt.Fprintln(stderr) + if err == nil { + return nil + } + var ee *exitError + if !errors.As(err, &ee) { + return err + } + } + } + fmt.Fprintln(stderr, loginRequiredMessage(cnf, lerr)) + return &exitError{code: auth.ExitCodeLoginRequired} +} + +// loginRequiredMessage matches the legacy CLI's LoginRequiredEvent::getExtendedMessage(). +func loginRequiredMessage(cnf *config.Config, lerr *auth.LoginRequiredError) string { + msg := lerr.Message() + if lerr.HasAPIToken { + if len(lerr.AuthMethods) == 1 && lerr.AuthMethods[0] == "mfa" { + msg += "\n\nThe API token may need to be re-created after enabling MFA." + } + return msg + } + loginCmd := "login" + if len(lerr.AuthMethods) > 0 { + loginCmd += " --method " + shellQuote(strings.Join(lerr.AuthMethods, ",")) + } + if lerr.MaxAge != nil { + loginCmd += fmt.Sprintf(" --max-age %d", *lerr.MaxAge) + } + return fmt.Sprintf("%s\n\nPlease log in by running:\n %s %s", msg, cnf.Application.Executable, loginCmd) +} + +// withLogin runs fn, and if it reports that login is required, offers a login and runs it again. +func withLogin(cmd *cobra.Command, cnf *config.Config, m *auth.Manager, fn func() error) error { + err := fn() + lerr, ok := auth.AsLoginRequired(err) + if !ok { + return err + } + if err := handleLoginRequired(cmd, cnf, m, lerr); err != nil { + return err + } + return fn() +} + +// hasDisplay matches the legacy CLI's Url::hasDisplay(). +func hasDisplay() bool { + if d := os.Getenv("DISPLAY"); d != "" { + return d != "none" + } + return runtime.GOOS == "windows" || runtime.GOOS == "darwin" +} + +// browserCommand returns the command to open URLs, or nil if none should be used. +func browserCommand(browserOption string) []string { + switch { + case browserOption == "0": + return nil + case strings.TrimSpace(browserOption) != "": + fields := strings.Fields(browserOption) + if _, err := exec.LookPath(fields[0]); err != nil { + return nil + } + return fields + case runtime.GOOS == "windows": + return []string{"rundll32", "url.dll,FileProtocolHandler"} + case runtime.GOOS == "darwin": + return []string{"open"} + } + for _, b := range []string{"xdg-open", "gnome-open"} { + if _, err := exec.LookPath(b); err == nil { + return []string{b} + } + } + return nil +} + +func canOpenURLs(browserOption string) bool { + return hasDisplay() && browserCommand(browserOption) != nil +} + +// openURL opens a URL in a browser, and reports whether it did. +func openURL(url, browserOption string) bool { + if !hasDisplay() { + debugLogf("Not opening URL (no display found)") + return false + } + args := browserCommand(browserOption) + if args == nil { + return false + } + //nolint:gosec // the browser is chosen by the user or the OS + return exec.Command(args[0], append(args[1:], url)...).Run() == nil +} + +func newAuthTokenCommand(cnf *config.Config) *cobra.Command { + cmd := &cobra.Command{ + Use: "auth:token", + Short: "Obtain an OAuth 2 access token for API requests", + Hidden: true, + Args: cobra.NoArgs, + Long: "This command prints a valid OAuth 2 access token to stdout. It can be used to make API requests via " + + "standard Bearer authentication (RFC 6750).\n\n" + + color.YellowString("Warning: access tokens must be kept secret.") + "\n\n" + + "Using this command is not generally recommended, as it increases the chance of the token being leaked. " + + "Take care not to expose the token in a shared program or system, or to send the token to the wrong " + + "API domain.", + Example: fmt.Sprintf(" # Print the payload for JWT-formatted tokens\n"+ + " %[1]s auth:token -W | cut -d. -f2 | base64 -d\n\n"+ + " # Use the token in a curl command\n curl -H\"$(%[1]s auth:token -HW)\" %[2]s/users/me", + cnf.Application.Executable, strings.TrimRight(cnf.API.BaseURL, "/")), + RunE: func(cmd *cobra.Command, _ []string) error { + if noWarn, _ := cmd.Flags().GetBool("no-warn"); !noWarn { + fmt.Fprintln(cmd.ErrOrStderr(), color.YellowString("Warning: keep access tokens secret.")) + } + m, err := newAuthManager(cnf, cmd.ErrOrStderr()) + if err != nil { + return err + } + var tok *auth.Token + if err := withLogin(cmd, cnf, m, func() (err error) { + tok, err = m.Token(cmd.Context(), "") + return err + }); err != nil { + return err + } + out := tok.AccessToken + if header, _ := cmd.Flags().GetBool("header"); header { + out = "Authorization: Bearer " + out + } + fmt.Fprint(cmd.OutOrStdout(), out) + return nil + }, + } + cmd.Flags().BoolP("header", "H", false, `Prefix the token with "Authorization: Bearer " to make an RFC 6750 header`) + cmd.Flags().BoolP("no-warn", "W", false, "Suppress the warning that is printed by default to stderr."+ + " This option is preferred over redirecting stderr, as that would hide other potentially useful messages.") + return cmd +} + +// newAuthInternalCommand is used by the legacy CLI to get tokens and auth state. It never prompts. +// +// Its stdout is only JSON. Exit code 3 means login is required, with the reason as JSON on the last line of stderr. +func newAuthInternalCommand(cnf *config.Config) *cobra.Command { + cmd := &cobra.Command{ + Use: "auth:internal", + Short: "Internal: provide tokens and auth state to the legacy CLI", + Hidden: true, + } + run := func(fn func(cmd *cobra.Command, m *auth.Manager) (any, error)) func(cmd *cobra.Command, _ []string) error { + return func(cmd *cobra.Command, _ []string) error { + m, err := newAuthManager(cnf, cmd.ErrOrStderr()) + if err != nil { + return err + } + result, err := fn(cmd, m) + if lerr, ok := auth.AsLoginRequired(err); ok { + b, _ := json.Marshal(lerr) + fmt.Fprintln(cmd.ErrOrStderr(), string(b)) + return &exitError{code: auth.ExitCodeLoginRequired} + } + if err != nil { + return err + } + return json.NewEncoder(cmd.OutOrStdout()).Encode(result) + } + } + tokenCmd := &cobra.Command{ + Use: "token", + Args: cobra.NoArgs, + RunE: run(func(cmd *cobra.Command, m *auth.Manager) (any, error) { + var rejected string + // The rejected token is read from stdin, to keep it out of process listings. + if r, _ := cmd.Flags().GetBool("rejected"); r { + b, err := io.ReadAll(cmd.InOrStdin()) + if err != nil { + return nil, err + } + rejected = strings.TrimSpace(string(b)) + } + return m.Token(cmd.Context(), rejected) + }), + } + tokenCmd.Flags().Bool("rejected", false, "Read an access token that was rejected by the API from stdin") + statusCmd := &cobra.Command{ + Use: "status", + Args: cobra.NoArgs, + RunE: run(func(cmd *cobra.Command, m *auth.Manager) (any, error) { + return m.Status(cmd.Context()) + }), + } + cmd.AddCommand(tokenCmd, statusCmd) + return cmd +} diff --git a/commands/auth_cleanup.go b/commands/auth_cleanup.go new file mode 100644 index 000000000..eb2fb6a41 --- /dev/null +++ b/commands/auth_cleanup.go @@ -0,0 +1,36 @@ +package commands + +import ( + "errors" + "os" + "path/filepath" + "slices" + + "github.com/upsun/cli/internal/config" +) + +// clearLegacySessionFiles deletes the legacy CLI's files for sessions that were logged out: the API cache, which can +// hold the previous account's data, each session's SSH certificate and config, and the current session's SSH include. +// With all set, the whole legacy session directory is deleted. +func clearLegacySessionFiles(cnf *config.Config, current string, ids []string, all bool) error { + dir, err := cnf.WritableUserDir() //nolint:staticcheck // the legacy CLI's files are in the user dir + if err != nil { + return err + } + sessionDir := filepath.Join(dir, ".session") + paths := []string{filepath.Join(dir, "cache")} + for _, id := range ids { + paths = append(paths, filepath.Join(sessionDir, "sess-cli-"+id)) + } + if all || slices.Contains(ids, current) { + paths = append(paths, filepath.Join(dir, "ssh", "session.config")) + } + if all { + paths = append(paths, sessionDir) + } + var errs []error + for _, p := range paths { + errs = append(errs, os.RemoveAll(p)) + } + return errors.Join(errs...) +} diff --git a/commands/auth_login.go b/commands/auth_login.go new file mode 100644 index 000000000..0c184c13c --- /dev/null +++ b/commands/auth_login.go @@ -0,0 +1,476 @@ +package commands + +import ( + "context" + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "fmt" + "html" + "net" + "net/http" + "net/url" + "regexp" + "strconv" + "strings" + "time" + + "github.com/fatih/color" + "github.com/spf13/cobra" + + "github.com/upsun/cli/internal/auth" + "github.com/upsun/cli/internal/auth/store" + "github.com/upsun/cli/internal/config" +) + +const loginTimeout = 30 * time.Minute + +type browserLoginOptions struct { + force bool + methods []string + maxAge *int + browser string + pipe bool +} + +func newBrowserLoginCommand(cnf *config.Config) *cobra.Command { + exe := cnf.Application.Executable + cmd := &cobra.Command{ + Use: "auth:browser-login", + Aliases: []string{"login"}, + Short: "Log in via a browser", + Args: cobra.NoArgs, + Long: fmt.Sprintf("Use this command to log in to the %s using a web browser.\n\n"+ + "It launches a temporary local website which redirects you to log in if necessary, "+ + "and then captures the resulting authorization code.\n\n"+ + "Your system's default browser will be used. You can override this using the --browser option.\n\n"+ + "Alternatively, to log in using an API token (without a browser), run: %s auth:api-token-login\n\n%s", + cnf.Application.Name, exe, nonInteractiveAuthHelp(cnf)), + RunE: func(cmd *cobra.Command, _ []string) error { + opts := &browserLoginOptions{} + opts.force, _ = cmd.Flags().GetBool("force") + methods, _ := cmd.Flags().GetStringSlice("method") + opts.methods = methods + if cmd.Flags().Changed("max-age") { + s, _ := cmd.Flags().GetString("max-age") + v, err := strconv.Atoi(s) + if err != nil || v < 0 { + fmt.Fprintln(cmd.ErrOrStderr(), "The --max-age value must be a non-negative integer.") + return &exitError{code: 1} + } + opts.maxAge = &v + } + opts.browser, _ = cmd.Flags().GetString("browser") + opts.pipe, _ = cmd.Flags().GetBool("pipe") + m, err := newAuthManager(cnf, cmd.ErrOrStderr()) + if err != nil { + return err + } + return runBrowserLogin(cmd, cnf, m, opts) + }, + } + cmd.Flags().BoolP("force", "f", false, "Log in again, even if already logged in") + cmd.Flags().StringSlice("method", nil, "Require specific authentication method(s)") + cmd.Flags().String("max-age", "", "The maximum age (in seconds) of the web authentication session") + cmd.Flags().String("browser", "", "The browser to use to open the URL. Set 0 for none.") + cmd.Flags().Bool("pipe", false, "Output the URL to stdout.") + return cmd +} + +func runBrowserLogin(cmd *cobra.Command, cnf *config.Config, m *auth.Manager, opts *browserLoginOptions) error { + ctx := cmd.Context() + stderr := cmd.ErrOrStderr() + if has, err := m.HasConfiguredToken(); err != nil { + return err + } else if has { + fmt.Fprintln(stderr, "Cannot log in via the browser, because an API token is set via config.") + return &exitError{code: 1} + } + if !isInteractive(cmd) { + fmt.Fprintln(stderr, "Non-interactive use of this command is not supported.") + fmt.Fprintln(stderr, "\n"+nonInteractiveAuthHelp(cnf)) + return &exitError{code: 1} + } + if advice := sessionAdvice(ctx, cnf, m, "Change this using"); advice != nil { + fmt.Fprintln(stderr, strings.Join(advice, "\n")) + fmt.Fprintln(stderr) + } + + if !opts.force && len(opts.methods) == 0 && opts.maxAge == nil { + status, err := m.Status(ctx) + if err != nil { + return err + } + if status.LoggedIn { + // Check whether the login is still valid. If so, only log in again if the user confirms. + account, err := getMyAccount(ctx, cnf, m) + if err == nil { + fmt.Fprintf(stderr, "You are already logged in as %s (%s)\n", + color.GreenString(account.Username), color.GreenString(account.Email)) + ok, err := confirm(cmd, "Log in anyway?", false) + if err != nil { + return err + } + if !ok { + return &exitError{code: 1} + } + opts.force = true + } else { + debugLogf("Already logged in, but a test request failed. Continuing with login: %s", err) + } + } + } + + // The system assigns a free port: the auth server allows any port for loopback redirects (RFC 8252). + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + return fmt.Errorf("failed to start a local server: %w", err) + } + localURL := "http://" + listener.Addr().String() + + verifier := randomString() + prompt := "consent" + if opts.force { + prompt = "consent select_account" + } + ls := &loginServer{ + cnf: cnf, + localURL: localURL, + authorize: m.Settings.AuthorizeURL, + clientID: m.Settings.ClientID, + state: randomString(), + challenge: pkceChallenge(verifier), + prompt: prompt, + methods: opts.methods, + maxAge: opts.maxAge, + result: make(chan loginResult, 1), + } + srv := &http.Server{Handler: ls, ReadHeaderTimeout: 10 * time.Second} + go func() { _ = srv.Serve(listener) }() + defer func() { + shutdownCtx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + _ = srv.Shutdown(shutdownCtx) + }() + + switch { + case opts.pipe: + fmt.Fprintln(cmd.OutOrStdout(), localURL) + case opts.browser != "0" && openURL(localURL, opts.browser): + fmt.Fprintf(stderr, "Opened URL: %s\n", color.GreenString(localURL)) + fmt.Fprintln(stderr, "Please use the browser to log in.") + default: + fmt.Fprintln(stderr, "Please open the following URL in a browser and log in:") + fmt.Fprintln(stderr, color.GreenString(localURL)) + } + fmt.Fprintln(stderr) + fmt.Fprintln(stderr, color.New(color.Bold).Sprint("Help:")) + fmt.Fprintln(stderr, " Leave this command running during login.") + fmt.Fprintln(stderr, " If you need to quit, use Ctrl+C.") + fmt.Fprintln(stderr) + + var res loginResult + select { + case res = <-ls.result: + case <-time.After(loginTimeout): + fmt.Fprintln(stderr, "Login timed out after 30 minutes") + fmt.Fprintln(stderr) + case <-ctx.Done(): + return ctx.Err() + } + // Allow a little time for the final page to be displayed in the browser. + time.Sleep(100 * time.Millisecond) + + if res.code == "" { + fmt.Fprintln(stderr, "Failed to get an authorization code.") + fmt.Fprintln(stderr) + switch { + case res.err != "" && res.errDescription != "": + fmt.Fprintln(stderr, " OAuth 2.0 error: "+color.RedString(res.err)) + fmt.Fprintln(stderr, " Description: "+res.errDescription) + if res.errHint != "" { + fmt.Fprintln(stderr, " Hint: "+res.errHint) + } + fmt.Fprintln(stderr) + case res.errDescription != "": + fmt.Fprintln(stderr, res.errDescription) + fmt.Fprintln(stderr) + } + fmt.Fprintln(stderr, "Please try again.") + return &exitError{code: 1} + } + + fmt.Fprintln(stderr, "Login information received. Verifying...") + entry, err := m.OAuth.ExchangeCode(ctx, res.code, verifier, localURL) + if err != nil { + return fmt.Errorf("failed to exchange the authorization code: %w", err) + } + if err := saveLogin(cmd, cnf, m, entry, ""); err != nil { + return err + } + + if entry.RefreshToken == "" { + fmt.Fprintln(stderr) + fmt.Fprintln(stderr, color.New(color.Bold, color.FgYellow).Sprint("Warning:")) + fmt.Fprintln(stderr, "No refresh token is available. This will cause frequent login errors.") + fmt.Fprintln(stderr, "Please contact support.") + fmt.Fprintf(stderr, "For internal use: the OAuth 2 client is probably misconfigured (client ID: %s).\n", + color.YellowString(m.Settings.ClientID)) + } + return nil +} + +// saveLogin logs out of the previous session, saves the new tokens, and runs the legacy CLI's post-login steps. +// For an API token login, entry holds the tokens from exchanging apiToken. +func saveLogin(cmd *cobra.Command, cnf *config.Config, m *auth.Manager, entry *store.Entry, apiToken string) error { + ctx := cmd.Context() + id := m.Settings.SessionID + if err := m.LogoutToReplace(ctx, id, apiToken); err != nil { + return err + } + if apiToken != "" { + if err := m.Save(ctx, auth.APITokenSessionID(apiToken), entry); err != nil { + return err + } + entry = &store.Entry{APIToken: apiToken} + } + if err := clearLegacySessionFiles(cnf, id, []string{id}, false); err != nil { + return err + } + if err := m.Save(ctx, id, entry); err != nil { + return err + } + stderr := cmd.ErrOrStderr() + fmt.Fprintln(stderr, "You are logged in.") + if err := runLegacyAuthHook(cmd, cnf, "auth:post-login"); err != nil { + return err + } + account, err := getMyAccount(ctx, cnf, m) + if err != nil { + return fmt.Errorf("failed to load account information: %w", err) + } + fmt.Fprintf(stderr, "\nUsername: %s\nEmail address: %s\n", + color.GreenString(account.Username), color.GreenString(account.Email)) + return nil +} + +type myAccount struct { + Username string `json:"username"` + Email string `json:"email"` +} + +// getMyAccount fetches the current user, without offering a login. +func getMyAccount(ctx context.Context, cnf *config.Config, m *auth.Manager) (*myAccount, error) { + u, err := url.JoinPath(m.Settings.BaseURL, "users", "me") + if err != nil { + return nil, err + } + req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, http.NoBody) + if err != nil { + return nil, err + } + resp, err := auth.NewClient(m, auth.NewHTTPClient(cnf, m.Settings).Transport).Do(req) + if err != nil { + return nil, err + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("unexpected status: %s", resp.Status) + } + var a myAccount + if err := json.NewDecoder(resp.Body).Decode(&a); err != nil { + return nil, err + } + return &a, nil +} + +func randomString() string { + b := make([]byte, 32) + _, _ = rand.Read(b) + return base64.RawURLEncoding.EncodeToString(b) +} + +// pkceChallenge applies the PKCE S256 transformation (RFC 7636). +func pkceChallenge(verifier string) string { + sum := sha256.Sum256([]byte(verifier)) + return base64.RawURLEncoding.EncodeToString(sum[:]) +} + +type loginResult struct { + code string + err, errDescription, errHint string +} + +// loginServer is the local web server that starts the OAuth 2.0 flow and receives the authorization code. +type loginServer struct { + cnf *config.Config + localURL string + authorize string + clientID string + state string + challenge string + prompt string + methods []string + maxAge *int + result chan loginResult +} + +type loginPage struct { + status int + location string + title string + content string +} + +func (s *loginServer) ServeHTTP(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/" { + http.NotFound(w, r) + return + } + p := s.handle(r.URL.Query()) + w.Header().Set("Cache-Control", "no-cache") + w.Header().Set("Content-Type", "text/html; charset=utf-8") + if p.location != "" { + w.Header().Set("Location", p.location) + } + w.WriteHeader(p.status) + _, _ = w.Write([]byte(s.render(p))) //nolint:gosec // values from the request are escaped +} + +func (s *loginServer) handle(q url.Values) *loginPage { + switch { + case q.Has("state") && q.Has("code"): + // The response after a successful OAuth 2.0 redirect. + if q.Get("state") != s.state { + return s.reportError("Invalid state parameter", "", "") + } + if q.Has("code_challenge") && q.Get("code_challenge") != s.challenge { + return s.reportError("Invalid returned code_challenge parameter", "", "") + } + s.send(loginResult{code: q.Get("code")}) + return &loginPage{ + status: http.StatusFound, + location: s.localURL + "/?done", + content: "

Authentication response received, please wait...

", + } + case q.Has("done"): + return &loginPage{ + status: http.StatusOK, + title: "Successfully logged in", + content: "

You can return to the command line

", + } + case q.Has("error"): + return s.reportError(q.Get("error_description"), q.Get("error"), q.Get("error_hint")) + } + authURL := s.authorizeURL() + return &loginPage{ + status: http.StatusFound, + location: authURL, + content: `

Log in.

`, + } +} + +// send reports a result to the command, once. +func (s *loginServer) send(r loginResult) { + select { + case s.result <- r: + default: + } +} + +func (s *loginServer) reportError(message, oauthErr, hint string) *loginPage { + p := &loginPage{status: http.StatusUnauthorized, title: "Error"} + if oauthErr != "" { + p.content += `

` + html.EscapeString(oauthErr) + `

` + } + if message != "" { + p.content += `

` + html.EscapeString(message) + `

` + } + if hint != "" { + p.content += `

` + html.EscapeString(hint) + `

` + } + if message != "" || oauthErr != "" || hint != "" { + s.send(loginResult{err: oauthErr, errDescription: message, errHint: hint}) + } + p.content += "

Please try again

" + return p +} + +func (s *loginServer) authorizeURL() string { + params := url.Values{ + "redirect_uri": {s.localURL}, + "state": {s.state}, + "client_id": {s.clientID}, + "prompt": {s.prompt}, + "response_type": {"code"}, + "code_challenge": {s.challenge}, + "code_challenge_method": {"S256"}, + "scope": {"offline_access"}, + } + if len(s.methods) > 0 { + params.Set("amr", strings.Join(s.methods, " ")) + } + if s.maxAge != nil { + params.Set("max_age", strconv.Itoa(*s.maxAge)) + } + sep := "?" + if strings.Contains(s.authorize, "?") { + sep = "&" + } + // PHP's http_build_query with RFC 3986 encoding uses %20 for spaces. + return s.authorize + sep + strings.ReplaceAll(params.Encode(), "+", "%20") +} + +var loginPlaceholder = regexp.MustCompile(`\{\{\s*(content|title)\s*}}`) + +func (s *loginServer) render(p *loginPage) string { + body := "

" + p.title + "

" + p.content + if tpl := s.cnf.BrowserLogin.Body; tpl != "" { + body = loginPlaceholder.ReplaceAllStringFunc(tpl, func(m string) string { + if strings.Contains(m, "content") { + return p.content + } + return p.title + }) + } + var css string + if s.cnf.BrowserLogin.CSS != "" { + css = "\n" + } + return ` + + + + ` + html.EscapeString(s.cnf.Application.Name) + `: Authentication (temporary URL) + + ` + css + ` + +` + body + ` + + +` +} diff --git a/commands/auth_logout.go b/commands/auth_logout.go new file mode 100644 index 000000000..bec225a06 --- /dev/null +++ b/commands/auth_logout.go @@ -0,0 +1,180 @@ +package commands + +import ( + "errors" + "fmt" + "slices" + "strings" + + "github.com/fatih/color" + "github.com/spf13/cobra" + + "github.com/upsun/cli/internal/auth" + "github.com/upsun/cli/internal/config" +) + +func newAPITokenLoginCommand(cnf *config.Config) *cobra.Command { + help := fmt.Sprintf("Use this command to log in to your %s account using an API token.", cnf.Service.Name) + help += fmt.Sprintf("\n\nAlternatively, to log in to the CLI with a browser, run:\n %s", + color.GreenString(cnf.Application.Executable+" auth:browser-login")) + return &cobra.Command{ + Use: "auth:api-token-login", + Short: "Log in using an API token", + Long: help, + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, _ []string) error { + stderr := cmd.ErrOrStderr() + m, err := newAuthManager(cnf, stderr) + if err != nil { + return err + } + if has, err := m.HasConfiguredToken(); err != nil { + return err + } else if has { + fmt.Fprintln(stderr, "An API token is already set via config") + return &exitError{code: 1} + } + if !isInteractive(cmd) { + fmt.Fprintln(stderr, "Non-interactive use of this command is not supported.") + fmt.Fprintln(stderr, "\n"+nonInteractiveAuthHelp(cnf)) + return &exitError{code: 1} + } + + const maxAttempts = 5 + for range maxAttempts { + fmt.Fprint(stderr, "Please enter an API token:\n> ") + apiToken, err := readSecret(cmd) + if err != nil { + return err + } + apiToken = strings.TrimSpace(apiToken) + if apiToken == "" { + fmt.Fprintln(stderr, color.RedString("The token cannot be empty")) + continue + } + entry, err := m.OAuth.ExchangeAPIToken(cmd.Context(), apiToken) + if err != nil { + var oerr *auth.OAuthError + if !errors.As(err, &oerr) { + return err + } + fmt.Fprintln(stderr, color.RedString(err.Error())) + continue + } + fmt.Fprintln(stderr) + fmt.Fprintln(stderr, "The API token is valid.") + return saveLogin(cmd, cnf, m, entry, apiToken) + } + // Each error has been printed. + return &exitError{code: 1} + }, + } +} + +func newLogoutCommand(cnf *config.Config) *cobra.Command { + cmd := &cobra.Command{ + Use: "auth:logout", + Aliases: []string{"logout"}, + Short: "Log out", + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, _ []string) error { + all, _ := cmd.Flags().GetBool("all") + other, _ := cmd.Flags().GetBool("other") + return runLogout(cmd, cnf, all, other) + }, + } + cmd.Flags().BoolP("all", "a", false, "Log out from all local sessions") + cmd.Flags().Bool("other", false, "Log out from other local sessions") + return cmd +} + +func runLogout(cmd *cobra.Command, cnf *config.Config, all, other bool) error { + ctx := cmd.Context() + stderr := cmd.ErrOrStderr() + m, err := newAuthManager(cnf, stderr) + if err != nil { + return err + } + + // API tokens set via config cannot be removed using this command. + if has, err := m.HasConfiguredToken(); err != nil { + return err + } else if has { + fmt.Fprintln(stderr, color.YellowString("Warning: an API token is set via config")) + } + + current := m.Settings.SessionID + ids, err := m.SessionIDs(ctx) + if err != nil { + return err + } + + if other && !all { + fmt.Fprintf(stderr, "The current session ID is: %s\n", color.GreenString(current)) + others := slices.DeleteFunc(slices.Clone(ids), func(id string) bool { return id == current }) + if len(others) == 0 { + fmt.Fprintln(stderr, "No other sessions exist.") + return nil + } + fmt.Fprintln(stderr) + for _, id := range others { + if err := m.Logout(ctx, id); err != nil { + return err + } + } + if err := clearLegacySessionFiles(cnf, current, others, false); err != nil { + return err + } + for _, id := range others { + fmt.Fprintf(stderr, "Logged out from session: %s\n", color.GreenString(id)) + } + fmt.Fprintln(stderr) + fmt.Fprintln(stderr, "All other sessions have been deleted.") + return nil + } + + if err := m.Logout(ctx, current); err != nil { + return err + } + if all { + for _, id := range ids { + if err := m.Logout(ctx, id); err != nil { + return err + } + } + if err := m.DeleteAll(ctx); err != nil { + return err + } + if err := clearLegacySessionFiles(cnf, current, ids, true); err != nil { + return err + } + fmt.Fprintln(stderr, "You are now logged out.") + fmt.Fprintln(stderr) + fmt.Fprintln(stderr, "All sessions have been deleted.") + printSessionAdvice(cmd, cnf, m) + return nil + } + if err := clearLegacySessionFiles(cnf, current, []string{current}, false); err != nil { + return err + } + fmt.Fprintln(stderr, "You are now logged out.") + printSessionAdvice(cmd, cnf, m) + + remaining, err := m.SessionIDs(ctx) + if err != nil { + return err + } + if len(remaining) > 0 { + fmt.Fprintln(stderr) + fmt.Fprintf(stderr, "Other sessions exist. Log out of all sessions with: %s\n", + color.YellowString(cnf.Application.Executable+" logout --all")) + } + return nil +} + +func printSessionAdvice(cmd *cobra.Command, cnf *config.Config, m *auth.Manager) { + if advice := sessionAdvice(cmd.Context(), cnf, m, "Change this using"); advice != nil { + fmt.Fprintln(cmd.ErrOrStderr()) + fmt.Fprintln(cmd.ErrOrStderr(), strings.Join(advice, "\n")) + } +} diff --git a/commands/auth_test.go b/commands/auth_test.go new file mode 100644 index 000000000..0704dfe44 --- /dev/null +++ b/commands/auth_test.go @@ -0,0 +1,38 @@ +package commands + +import ( + "testing" + + "github.com/platformsh/platformify/vendorization" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestAuthCommands_Disabled(t *testing.T) { + cnf := testConfig() + cnf.Application.DisabledCommands = []string{"auth:api-token-login"} + cnf.Application.WrappedDisabledCommands = []string{"auth:token"} + names := make([]string, 0, 2) + for _, c := range authCommands(cnf) { + names = append(names, c.Name()) + } + assert.Equal(t, []string{"auth:browser-login", "auth:logout"}, names) +} + +func TestAuthCommands_LegacyGlobalFlags(t *testing.T) { + root := newRootCommand(testConfig(), &vendorization.VendorAssets{}) + for _, args := range [][]string{ + {"logout", "-n"}, + {"auth:token", "-W", "--no-ansi"}, + {"login", "--ansi", "--no"}, + } { + c, rest, err := root.Find(args) + require.NoError(t, err) + assert.NoError(t, c.ParseFlags(rest), "args: %v", args) + } +} + +func TestBrowserCommand_Whitespace(t *testing.T) { + assert.Nil(t, browserCommand("0")) + assert.NotPanics(t, func() { browserCommand(" ") }) +} diff --git a/commands/init.go b/commands/init.go index f2fab75fc..0cb85f7cb 100644 --- a/commands/init.go +++ b/commands/init.go @@ -5,6 +5,7 @@ import ( "context" "fmt" "io" + "net/http" "os" "path/filepath" "strings" @@ -108,11 +109,14 @@ func runInitCommand( cnf := config.FromContext(cmd.Context()) - legacyCLIClient, err := auth.NewLegacyCLIClient(cmd.Context(), - makeLegacyCLIWrapper(cnf, cmd.OutOrStdout(), cmd.ErrOrStderr(), cmd.InOrStdin())) + authManager, err := newAuthManager(cnf, cmd.ErrOrStderr()) if err != nil { return err } + httpClient := auth.NewClient(authManager, auth.NewHTTPClient(cnf, authManager.Settings).Transport) + ensureAuthenticated := func() error { + return withLogin(cmd, cnf, authManager, func() error { return authManager.EnsureAuthenticated(cmd.Context()) }) + } msg, canUse := canUseAI(cnf) if !canUse { @@ -130,7 +134,7 @@ func runInitCommand( var isInteractive = !viper.GetBool("no-interaction") debugLogf("Checking selected organization") - org, err := handleOrganizations(cmd.Context(), cnf, legacyCLIClient, initOptions) + org, err := handleOrganizations(cmd.Context(), cnf, httpClient, ensureAuthenticated, initOptions) if err != nil { return err } @@ -168,7 +172,7 @@ func runInitCommand( "Note: AI configuration is only compatible with `%s` organizations\n", api.OrgTypeFlexible)) } - if err := legacyCLIClient.EnsureAuthenticated(cmd.Context()); err != nil { + if err := ensureAuthenticated(); err != nil { return err } @@ -178,8 +182,8 @@ func runInitCommand( return err } - initOptions.HTTPClient = legacyCLIClient.HTTPClient - initOptions.APIURL = cnf.API.BaseURL + initOptions.HTTPClient = httpClient + initOptions.APIURL = authManager.Settings.BaseURL initOptions.UserAgent = cnf.UserAgent() initOptions.IsInteractive = isInteractive initOptions.Yes = viper.GetBool("yes") @@ -192,13 +196,17 @@ func runInitCommand( // handleOrganizations manages organization selection and validation. // It modifies initOptions.OrganizationID and initOptions.ProjectID. func handleOrganizations( - ctx context.Context, cnf *config.Config, legacyCLIClient *auth.LegacyCLIClient, initOptions *_init.Options, + ctx context.Context, + cnf *config.Config, + httpClient *http.Client, + ensureAuthenticated func() error, + initOptions *_init.Options, ) (*api.Organization, error) { if !cnf.API.EnableOrganizations { return nil, nil } - apiClient, err := api.NewClient(cnf.API.BaseURL, legacyCLIClient.HTTPClient) + apiClient, err := api.NewClient(cnf.API.BaseURL, httpClient) if err != nil { return nil, err } @@ -210,7 +218,7 @@ func handleOrganizations( return nil, nil } - if err := legacyCLIClient.EnsureAuthenticated(ctx); err != nil { + if err := ensureAuthenticated(); err != nil { return nil, err } diff --git a/commands/list.go b/commands/list.go index 755db330b..56b2c1960 100644 --- a/commands/list.go +++ b/commands/list.go @@ -60,6 +60,17 @@ func newListCommand(cnf *config.Config) *cobra.Command { list.AddCommand(&appProjectConvertCommand) } + for _, c := range authCommands(cnf) { + desc := commandFromCobra(cnf, c) + if desc.Hidden && !viper.GetBool("all") { + continue + } + if !list.DescribesNamespace() || list.Namespace == desc.Name.Namespace { + list.RemoveCommand(desc.Name.String()) + list.AddCommand(&desc) + } + } + format := viper.GetString("format") raw := viper.GetBool("raw") diff --git a/commands/list_cobra.go b/commands/list_cobra.go new file mode 100644 index 000000000..60e228f97 --- /dev/null +++ b/commands/list_cobra.go @@ -0,0 +1,75 @@ +package commands + +import ( + "fmt" + "strings" + + "github.com/fatih/color" + "github.com/spf13/cobra" + "github.com/spf13/pflag" + orderedmap "github.com/wk8/go-ordered-map/v2" + + "github.com/upsun/cli/internal/config" +) + +// commandFromCobra describes a native Cobra command in the same format as the legacy CLI's commands. +func commandFromCobra(cnf *config.Config, c *cobra.Command) Command { + namespace, name, ok := strings.Cut(c.Name(), ":") + if !ok { + namespace, name = "", c.Name() + } + options := orderedmap.New[string, Option]() + c.LocalNonPersistentFlags().VisitAll(func(f *pflag.Flag) { + if f.Hidden { + return + } + opt := Option{ + Name: "--" + f.Name, + Description: CleanString(f.Usage), + } + if f.Shorthand != "" { + opt.Shortcut = "-" + f.Shorthand + } + switch f.Value.Type() { + case "bool": + opt.Default = Any{false} + case "stringSlice": + opt.AcceptValue, opt.IsValueRequired, opt.IsMultiple = true, true, true + opt.Default = Any{[]any{}} + default: + opt.AcceptValue, opt.IsValueRequired = true, true + opt.Default = Any{nil} + } + options.Set(f.Name, opt) + }) + for _, opt := range globalOptions(cnf) { + if !opt.Hidden { + options.Set(opt.GetName(), opt) + } + } + return Command{ + Name: CommandName{Namespace: namespace, Command: name}, + Usage: []string{fmt.Sprintf("%s %s", cnf.Application.Executable, c.Name())}, + Aliases: c.Aliases, + Description: CleanString(c.Short), + Help: CleanString(c.Long), + Definition: Definition{ + Arguments: orderedmap.New[string, Argument](), + Options: options, + }, + Hidden: c.Hidden, + } +} + +// useLegacyStyleHelp makes a native command print its help in the same format as the legacy CLI's commands. +func useLegacyStyleHelp(cnf *config.Config, c *cobra.Command) *cobra.Command { + c.SetHelpFunc(func(c *cobra.Command, _ []string) { + desc := commandFromCobra(cnf, c) + fmt.Fprintln(c.OutOrStdout(), desc.HelpPage(cnf)) + if c.Example != "" { + fmt.Fprintln(c.OutOrStdout(), color.YellowString("Examples:")) + fmt.Fprintln(c.OutOrStdout(), c.Example) + } + }) + return c +} diff --git a/commands/list_models.go b/commands/list_models.go index 690b249ef..a7a9dba46 100644 --- a/commands/list_models.go +++ b/commands/list_models.go @@ -6,6 +6,7 @@ import ( "fmt" "math" "regexp" + "slices" "sort" "strings" "text/tabwriter" @@ -533,3 +534,11 @@ func (l *List) AddCommand(cmd *Command) { } }) } + +// RemoveCommand removes a command by name, e.g. a legacy command that is replaced by a native one. +func (l *List) RemoveCommand(name string) { + l.Commands = slices.DeleteFunc(l.Commands, func(c *Command) bool { return c.Name.String() == name }) + for i := range l.Namespaces { + l.Namespaces[i].Commands = slices.DeleteFunc(l.Namespaces[i].Commands, func(n string) bool { return n == name }) + } +} diff --git a/commands/root.go b/commands/root.go index 0db2a8240..247b4172c 100644 --- a/commands/root.go +++ b/commands/root.go @@ -42,7 +42,20 @@ func Execute(cnf *config.Config) error { } else if ok { cmd.SetArgs(args) } - return cmd.ExecuteContext(ctx) + err := cmd.ExecuteContext(ctx) + var ee *exitError + if errors.As(err, &ee) { + os.Exit(ee.code) + } + if err != nil && !isQuiet() { + fmt.Fprintln(color.Error, "Error:", err) + } + return err +} + +// isQuiet reports whether quiet mode is on, which --debug and --verbose override. +func isQuiet() bool { + return viper.GetBool("quiet") && !viper.GetBool("debug") && !viper.GetBool("verbose") } func newRootCommand(cnf *config.Config, assets *vendorization.VendorAssets) *cobra.Command { @@ -54,13 +67,14 @@ func newRootCommand(cnf *config.Config, assets *vendorization.VendorAssets) *cob DisableFlagParsing: false, FParseErrWhitelist: cobra.FParseErrWhitelist{UnknownFlags: true}, SilenceUsage: true, - SilenceErrors: false, + // Errors are printed by Execute, which handles exit codes. + SilenceErrors: true, PersistentPreRun: func(cmd *cobra.Command, _ []string) { - if isCompletionRequest(cmd) { - // Completions must be fast and quiet. + if isCompletionRequest(cmd) || isInternalCommand(cmd) { + // Completions and internal commands must be fast and quiet. return } - quiet := viper.GetBool("quiet") && !viper.GetBool("debug") && !viper.GetBool("verbose") + quiet := isQuiet() if quiet { viper.Set("no-interaction", true) cmd.SetErr(io.Discard) @@ -112,7 +126,7 @@ func newRootCommand(cnf *config.Config, assets *vendorization.VendorAssets) *cob } }, PersistentPostRun: func(cmd *cobra.Command, _ []string) { - if isCompletionRequest(cmd) { + if isCompletionRequest(cmd) || isInternalCommand(cmd) { return } checkShellConfigLeftovers(cmd.ErrOrStderr(), cnf) @@ -160,6 +174,7 @@ func newRootCommand(cnf *config.Config, assets *vendorization.VendorAssets) *cob // Add subcommands. cmd.AddCommand( + newAuthInternalCommand(cnf), newCompleteCommand(cnf), newConfigInstallCommand(), newCompletionCommand(cnf), @@ -172,6 +187,11 @@ func newRootCommand(cnf *config.Config, assets *vendorization.VendorAssets) *cob if cnf.Service.ProjectConfigFlavor == "upsun" { cmd.AddCommand(newProjectConvertCommand(cnf)) } + for _, c := range authCommands(cnf) { + addLegacyGlobalFlags(c) + c.PreRun = func(c *cobra.Command, _ []string) { applyLegacyGlobalFlags(c) } + cmd.AddCommand(useLegacyStyleHelp(cnf, c)) + } // Define the help flag before Cobra looks up the command, so that "--help init" does not treat "init" as its value. cmd.InitDefaultHelpFlag() @@ -182,6 +202,16 @@ func newRootCommand(cnf *config.Config, assets *vendorization.VendorAssets) *cob return cmd } +// isInternalCommand reports whether the command is used internally by the legacy CLI. +func isInternalCommand(cmd *cobra.Command) bool { + for c := cmd; c != nil; c = c.Parent() { + if c.Name() == "auth:internal" { + return true + } + } + return false +} + // checkShellConfigLeftovers checks .zshrc and .bashrc for any leftovers from the legacy CLI func checkShellConfigLeftovers(w io.Writer, cnf *config.Config) { start := fmt.Sprintf("# BEGIN SNIPPET: %s configuration", cnf.Application.Name) @@ -301,7 +331,7 @@ func exitWithError(err error) { debugLogf(err.Error()) os.Exit(exitCode) } - if !viper.GetBool("quiet") { + if !isQuiet() { fmt.Fprintln(color.Error, color.RedString(err.Error())) } os.Exit(1) diff --git a/go.mod b/go.mod index c398f9e74..e4a74fc59 100644 --- a/go.mod +++ b/go.mod @@ -10,6 +10,7 @@ require ( github.com/fatih/color v1.19.0 github.com/go-chi/chi/v5 v5.3.2 github.com/go-playground/validator/v10 v10.30.5 + github.com/godbus/dbus/v5 v5.2.2 github.com/gofrs/flock v0.13.1 github.com/oklog/ulid/v2 v2.1.2 github.com/platformsh/platformify v0.5.0 @@ -21,8 +22,8 @@ require ( github.com/upsun/lib-sun v0.3.16 github.com/upsun/whatsun v0.2.1 github.com/wk8/go-ordered-map/v2 v2.1.8 + github.com/zalando/go-keyring v0.2.8 golang.org/x/crypto v0.57.0 - golang.org/x/oauth2 v0.37.0 golang.org/x/sync v0.23.0 golang.org/x/sys v0.48.0 golang.org/x/term v0.46.0 @@ -57,6 +58,7 @@ require ( github.com/clipperhouse/uax29/v2 v2.7.0 // indirect github.com/cloudflare/circl v1.6.3 // indirect github.com/cyphar/filepath-securejoin v0.6.1 // indirect + github.com/danieljoos/wincred v1.2.3 // indirect github.com/dlclark/regexp2/v2 v2.2.1 // indirect github.com/dsnet/compress v0.0.2-0.20230904184137-39efe44ab707 // indirect github.com/emirpasic/gods v1.18.1 // indirect diff --git a/go.sum b/go.sum index a4dfbaf6f..cd0967398 100644 --- a/go.sum +++ b/go.sum @@ -74,6 +74,8 @@ github.com/creack/pty v1.1.17 h1:QeVUsEDNrLBW4tMgZHvxy18sKtr6VI492kBhUfhDJNI= github.com/creack/pty v1.1.17/go.mod h1:MOBLtS5ELjhRRrroQr9kyvTxUAFNvYEK993ew/Vr4O4= github.com/cyphar/filepath-securejoin v0.6.1 h1:5CeZ1jPXEiYt3+Z6zqprSAgSWiggmpVyciv8syjIpVE= github.com/cyphar/filepath-securejoin v0.6.1/go.mod h1:A8hd4EnAeyujCJRrICiOWqjS1AX0a9kM5XL+NwKoYSc= +github.com/danieljoos/wincred v1.2.3 h1:v7dZC2x32Ut3nEfRH+vhoZGvN72+dQ/snVXo/vMFLdQ= +github.com/danieljoos/wincred v1.2.3/go.mod h1:6qqX0WNrS4RzPZ1tnroDzq9kY3fu1KwE7MRLQK4X0bs= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/dlclark/regexp2/v2 v2.2.1 h1:mf4KkFUj0gJuarK8P+LgiS+Lit7m9N1yAwEfPbee7R0= @@ -123,6 +125,8 @@ github.com/go-playground/validator/v10 v10.30.5 h1:YyCXvVShZbs2Sm3Mb53eNOlhRXctS github.com/go-playground/validator/v10 v10.30.5/go.mod h1:wEqiaov48pXX1kjhc3Da8y0M0Dtg/BK7gurFBLgwFrQ= github.com/go-viper/mapstructure/v2 v2.5.0 h1:vM5IJoUAy3d7zRSVtIwQgBj7BiWtMPfmPEgAXnvj1Ro= github.com/go-viper/mapstructure/v2 v2.5.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= +github.com/godbus/dbus/v5 v5.2.2 h1:TUR3TgtSVDmjiXOgAAyaZbYmIeP3DPkld3jgKGV8mXQ= +github.com/godbus/dbus/v5 v5.2.2/go.mod h1:3AAv2+hPq5rdnr5txxxRwiGjPXamgoIHgz9FPBfOp3c= github.com/gofrs/flock v0.13.1 h1:jjREztyBeSKBZYAC+mgc1laB+xsgy4kYMf3FbKF2UBo= github.com/gofrs/flock v0.13.1/go.mod h1:sf4BFiHwnvgxa25DlQoDqXQnwRMEOwqxRq37P6MzzmE= github.com/golang/groupcache v0.0.0-20241129210726-2c02b8208cf8 h1:f+oWsMOmNPc8JmEHVZIycC7hBoQxHH9pNKQORJNozsQ= @@ -302,6 +306,8 @@ github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e/go.mod h1:RbqR21r5mrJu github.com/xyproto/randomstring v1.0.5 h1:YtlWPoRdgMu3NZtP45drfy1GKoojuR7hmRcnhZqKjWU= github.com/xyproto/randomstring v1.0.5/go.mod h1:rgmS5DeNXLivK7YprL0pY+lTuhNQW3iGxZ18UQApw/E= github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY= +github.com/zalando/go-keyring v0.2.8 h1:6sD/Ucpl7jNq10rM2pgqTs0sZ9V3qMrqfIIy5YPccHs= +github.com/zalando/go-keyring v0.2.8/go.mod h1:tsMo+VpRq5NGyKfxoBVjCuMrG47yj8cmakZDO5QGii0= github.com/zricethezav/gitleaks/v8 v8.30.1 h1:PmEvCfVI7ti9dV3s5aMZUY7sS2GxRvG3yzih7E+cS3w= github.com/zricethezav/gitleaks/v8 v8.30.1/go.mod h1:rTDwxRjufMKAkhTI/Mijd07nday1yOhf9qywjwz5Irw= go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= @@ -327,8 +333,6 @@ golang.org/x/net v0.0.0-20211112202133-69e39bad7dc2/go.mod h1:9nx3DQGgdP8bBQD5qx golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c= golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To= golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU= -golang.org/x/oauth2 v0.37.0 h1:JUlcxA8oAtauLfiH8FX2/FkAWHAdi0QtGCGc+hofE98= -golang.org/x/oauth2 v0.37.0/go.mod h1:IxwZNxUULJmpBFf9K/9NTMSIfZZuvuTy1gGxhigP/58= golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk= diff --git a/integration-tests/auth_browser_login_test.go b/integration-tests/auth_browser_login_test.go index c062a6f58..31af92e9d 100644 --- a/integration-tests/auth_browser_login_test.go +++ b/integration-tests/auth_browser_login_test.go @@ -130,8 +130,8 @@ func TestAuthBrowserLogin_Success(t *testing.T) { require.NoError(t, err) } -// writeOAuthSession writes a pre-populated OAuth session directly to the filesystem for a given -// homeDir and session ID. This bypasses the session.Manager so integration tests can set up +// writeOAuthSession writes an OAuth session in the legacy CLI's file format, for a given homeDir and +// session ID. The CLI migrates it to its own storage on first use, so integration tests can set up // authenticated state without running a full login flow. func writeOAuthSession(t *testing.T, homeDir, sessionID string, s map[string]any) { t.Helper() diff --git a/integration-tests/auth_flows_test.go b/integration-tests/auth_flows_test.go new file mode 100644 index 000000000..d041d43c5 --- /dev/null +++ b/integration-tests/auth_flows_test.go @@ -0,0 +1,201 @@ +package tests + +import ( + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/upsun/cli/pkg/mockapi" +) + +// newInteractiveLoginFactory returns a command factory for a user who is not logged in, in an interactive terminal +// with a browser that completes the login. +func newInteractiveLoginFactory(t *testing.T, apiURL, authURL string) *cmdFactory { + f := newCommandFactory(t, apiURL, authURL) + f.extraEnv = append(f.extraEnv, + EnvPrefix+"TOKEN=", + EnvPrefix+"NO_INTERACTION=", + "SHELL_INTERACTIVE=1", + // Do not ask whether to write ~/.ssh/config. + EnvPrefix+"API_WRITE_USER_SSH_CONFIG=0", + ) + f.fakeBrowserThatLogsIn() + return f +} + +func newMyUserAPI(t *testing.T) *httptest.Server { + apiHandler := mockapi.NewHandler(t) + apiHandler.SetMyUser(&mockapi.User{ID: "u1", Username: "testuser", Email: "test@example.com"}) + return httptest.NewServer(apiHandler) +} + +// TestAuthLogin_AcceptPromptInPHP accepts the login prompt in a PHP command, which runs the Go login. +func TestAuthLogin_AcceptPromptInPHP(t *testing.T) { + authServer := mockapi.NewAuthServer(t) + defer authServer.Close() + apiServer := newMyUserAPI(t) + defer apiServer.Close() + + f := newInteractiveLoginFactory(t, apiServer.URL, authServer.URL) + f.stdin = strings.NewReader("y\n") + stdout, stderr, err := f.RunCombinedOutput("auth:info", "id") + require.NoError(t, err, "stderr: %s", stderr) + assert.Equal(t, "u1\n", stdout) + assert.Contains(t, stderr, "Log in via a browser?") + assert.Contains(t, stderr, "You are logged in.") + + f.stdin = nil + assert.Equal(t, "access-token-1", f.Run("auth:token", "--no-warn")) +} + +// TestAuthLogin_AcceptPromptInGo accepts the login prompt in a Go command. +func TestAuthLogin_AcceptPromptInGo(t *testing.T) { + authServer := mockapi.NewAuthServer(t) + defer authServer.Close() + apiServer := newMyUserAPI(t) + defer apiServer.Close() + + f := newInteractiveLoginFactory(t, apiServer.URL, authServer.URL) + f.stdin = strings.NewReader("y\n") + stdout, stderr, err := f.RunCombinedOutput("auth:token", "--no-warn") + require.NoError(t, err, "stderr: %s", stderr) + assert.Equal(t, "access-token-1", stdout) + assert.Contains(t, stderr, "Log in via a browser?") +} + +// TestAuthBrowserLogin_ReplacesSession checks the effects of a login over an existing session. +func TestAuthBrowserLogin_ReplacesSession(t *testing.T) { + authServer := mockapi.NewAuthServer(t) + defer authServer.Close() + apiServer := newMyUserAPI(t) + defer apiServer.Close() + + f := newInteractiveLoginFactory(t, apiServer.URL, authServer.URL) + writeOAuthSession(t, f.home, "default", map[string]any{ + "accessToken": "old-access-token", + "expires": time.Now().Add(time.Hour).Unix(), + "refreshToken": "old-refresh-token", + }) + + _, stderr, err := f.RunCombinedOutput("auth:browser-login", "--force") + require.NoError(t, err, "stderr: %s", stderr) + assert.Contains(t, stderr, "You are logged in.") + assert.Contains(t, stderr, "Username: testuser\nEmail address: test@example.com") + + // The previous session was revoked and replaced. + assert.Contains(t, authServer.RevokedTokens(), "old-access-token") + assert.Contains(t, authServer.RevokedTokens(), "old-refresh-token") + assert.FileExists(t, filepath.Join(f.home, ".platform-test-cli", "auth", "default.json")) + assert.Equal(t, "access-token-1", f.Run("auth:token", "--no-warn")) + + // An SSH certificate was generated. + assert.FileExists(t, filepath.Join(f.home, ".platform-test-cli", ".session", "sess-cli-default", "ssh", + "id_ed25519-cert.pub")) +} + +// TestAuthRefresh_RejectedToken checks that PHP gets a new token from Go when the API rejects one before it expires. +func TestAuthRefresh_RejectedToken(t *testing.T) { + authServer := mockapi.NewAuthServer(t) + defer authServer.Close() + authServer.SetUniqueAccessTokens(true) + authServer.AddRefreshToken("initial-refresh-token") + apiHandler := mockapi.NewHandler(t) + apiHandler.SetMyUser(&mockapi.User{ID: "u1"}) + apiHandler.RejectAccessToken("revoked-token") + apiServer := httptest.NewServer(apiHandler) + defer apiServer.Close() + + f := newCommandFactory(t, apiServer.URL, authServer.URL) + f.extraEnv = append(f.extraEnv, EnvPrefix+"TOKEN=") + writeOAuthSession(t, f.home, "default", map[string]any{ + "accessToken": "revoked-token", + "expires": time.Now().Add(time.Hour).Unix(), + "refreshToken": "initial-refresh-token", + }) + + assert.Equal(t, "u1\n", f.Run("auth:info", "id", "--refresh")) + assert.Equal(t, 1, authServer.RefreshRequests()) + assert.Equal(t, "access-token-1", f.Run("auth:token", "--no-warn")) + assert.False(t, authServer.ReuseDetected()) +} + +// TestAuthLogout_AllRevokesEverySession checks the effects of logging out of all sessions. +func TestAuthLogout_AllRevokesEverySession(t *testing.T) { + authServer := mockapi.NewAuthServer(t) + defer authServer.Close() + + f := newCommandFactory(t, "", authServer.URL) + f.extraEnv = append(f.extraEnv, EnvPrefix+"TOKEN=") + future := time.Now().Add(time.Hour).Unix() + for _, id := range []string{"default", "other"} { + writeOAuthSession(t, f.home, id, map[string]any{ + "accessToken": id + "-access", "refreshToken": id + "-refresh", "expires": future, + }) + } + // Create lock files, as a refresh would. + assert.Equal(t, "default-access", f.Run("auth:token", "--no-warn")) + + _, stderr, err := f.RunCombinedOutput("auth:logout", "--all") + require.NoError(t, err, "stderr: %s", stderr) + assert.Contains(t, stderr, "All sessions have been deleted.") + assert.ElementsMatch(t, + []string{"default-access", "default-refresh", "other-access", "other-refresh"}, authServer.RevokedTokens()) + + // Only the migration marker and lock files are kept. + entries, err := os.ReadDir(filepath.Join(f.home, ".platform-test-cli", "auth")) + require.NoError(t, err) + for _, e := range entries { + assert.True(t, e.Name() == ".migrated" || strings.HasSuffix(e.Name(), ".lock"), "unexpected file: %s", e.Name()) + } + assert.NoDirExists(t, filepath.Join(f.home, ".platform-test-cli", ".session")) + + _, _, err = f.RunCombinedOutput("auth:token", "--no-warn") + assertExitCode(t, 3, err) +} + +// TestAuthAPITokenLogin_Logout checks that logging out deletes a stored API token and revokes its tokens. +func TestAuthAPITokenLogin_Logout(t *testing.T) { + authServer := mockapi.NewAuthServer(t) + defer authServer.Close() + apiServer := newMyUserAPI(t) + defer apiServer.Close() + + f := newCommandFactory(t, apiServer.URL, authServer.URL) + f.extraEnv = append(f.extraEnv, EnvPrefix+"TOKEN=", EnvPrefix+"API_WRITE_USER_SSH_CONFIG=0") + _, stderr, err := f.RunInteractive(mockapi.ValidAPITokens[0]+"\n", "auth:api-token-login") + require.NoError(t, err, "stderr: %s", stderr) + assert.JSONEq(t, `{"logged_in": true, "session_ids": ["default"], "has_stored_api_token": true}`, + f.Run("auth:internal", "status")) + + f.Run("auth:logout") + assert.Contains(t, authServer.RevokedTokens(), "refresh-token-1") + assert.JSONEq(t, `{"logged_in": false, "session_ids": [], "has_stored_api_token": false}`, + f.Run("auth:internal", "status")) +} + +// TestSessionSwitch_UsesNewSession checks that PHP and Go use the same session after a switch. +func TestSessionSwitch_UsesNewSession(t *testing.T) { + authServer := mockapi.NewAuthServer(t) + defer authServer.Close() + apiServer := newMyUserAPI(t) + defer apiServer.Close() + + f := newCommandFactory(t, apiServer.URL, authServer.URL) + f.extraEnv = append(f.extraEnv, EnvPrefix+"TOKEN=") + writeOAuthSession(t, f.home, "work", map[string]any{ + "accessToken": "work-token", "expires": time.Now().Add(time.Hour).Unix(), + }) + + // The switch reports the new session's account, which PHP gets with a token from Go. + _, stderr, err := f.RunCombinedOutput("session:switch", "work") + require.NoError(t, err, "stderr: %s", stderr) + assert.Contains(t, stderr, "Username: testuser") + + assert.Equal(t, "work-token", f.Run("auth:token", "--no-warn")) +} diff --git a/integration-tests/auth_go_test.go b/integration-tests/auth_go_test.go new file mode 100644 index 000000000..ffdac3ab9 --- /dev/null +++ b/integration-tests/auth_go_test.go @@ -0,0 +1,313 @@ +package tests + +import ( + "encoding/json" + "net/http/httptest" + "os" + "path/filepath" + "runtime" + "strings" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/upsun/cli/pkg/mockapi" +) + +// TestAuthMigration checks that sessions saved by the legacy CLI are migrated, and the legacy copies deleted. +func TestAuthMigration(t *testing.T) { + authServer := mockapi.NewAuthServer(t) + defer authServer.Close() + authServer.AddRefreshToken("legacy-refresh-token") + apiHandler := mockapi.NewHandler(t) + apiHandler.SetMyUser(&mockapi.User{ID: "u1"}) + apiServer := httptest.NewServer(apiHandler) + defer apiServer.Close() + + f := newCommandFactory(t, apiServer.URL, authServer.URL) + f.extraEnv = append(f.extraEnv, EnvPrefix+"TOKEN=") + + writeOAuthSession(t, f.home, "default", map[string]any{ + "accessToken": "legacy-access-token", + "tokenType": "bearer", + "expires": time.Now().Add(-time.Hour).Unix(), + "refreshToken": "legacy-refresh-token", + }) + sessDir := filepath.Join(f.home, ".platform-test-cli", ".session") + sshFile := filepath.Join(sessDir, "sess-cli-default", "ssh", "id_ed25519-cert.pub") + require.NoError(t, os.WriteFile(sshFile, []byte("cert"), 0o600)) + // An API token saved by auth:api-token-login in the "work" session. + require.NoError(t, os.MkdirAll(filepath.Join(sessDir, "sess-cli-work"), 0o700)) + apiTokenFile := filepath.Join(sessDir, "sess-cli-work", "api-token") + require.NoError(t, os.WriteFile(apiTokenFile, []byte(mockapi.ValidAPITokens[0]), 0o600)) + + assert.Equal(t, "access-token-1", f.Run("auth:token", "--no-warn")) + assert.False(t, authServer.ReuseDetected()) + + authDir := filepath.Join(f.home, ".platform-test-cli", "auth") + assert.FileExists(t, filepath.Join(authDir, ".migrated")) + assert.FileExists(t, filepath.Join(authDir, "default.json")) + assert.FileExists(t, filepath.Join(authDir, "work.json")) + assert.NoDirExists(t, filepath.Join(sessDir, "sess-default")) + assert.NoFileExists(t, apiTokenFile) + assert.FileExists(t, sshFile, "SSH certificates must be kept") + + // The migrated API token is used in its session. + f.extraEnv = append(f.extraEnv, EnvPrefix+"SESSION_ID=work") + assert.Equal(t, "access-token-1", f.Run("auth:token", "--no-warn")) + + // A legacy session written after the migration is not imported. + writeOAuthSession(t, f.home, "late", map[string]any{ + "accessToken": "late-token", + "expires": time.Now().Add(time.Hour).Unix(), + }) + f.extraEnv = append(f.extraEnv, EnvPrefix+"SESSION_ID=late") + _, _, err := f.RunCombinedOutput("auth:token", "--no-warn") + assertExitCode(t, 3, err) +} + +// TestAuthRefresh_ManyExpiries runs Go and PHP commands in parallel across several token expiries. +// Each expiry must cause exactly one refresh, without reusing a rotated refresh token. +func TestAuthRefresh_ManyExpiries(t *testing.T) { + authServer := mockapi.NewAuthServer(t) + defer authServer.Close() + // Tokens are refreshed 2 minutes before they expire, so they are usable for 3 seconds. + authServer.SetTokenLifetime(123 * time.Second) + authServer.AddRefreshToken("initial-refresh-token") + + apiHandler := mockapi.NewHandler(t) + apiHandler.SetMyUser(&mockapi.User{ID: "u1"}) + apiServer := httptest.NewServer(apiHandler) + defer apiServer.Close() + + f := newCommandFactory(t, apiServer.URL, authServer.URL) + f.extraEnv = append(f.extraEnv, EnvPrefix+"TOKEN=") + f.dir = t.TempDir() + getCommandName(t) + writeOAuthSession(t, f.home, "default", map[string]any{ + "accessToken": "expired-token", + "tokenType": "bearer", + "expires": time.Now().Add(-time.Hour).Unix(), + "refreshToken": "initial-refresh-token", + }) + // Migrate before running commands concurrently. + f.Run("auth:token", "--no-warn") + + const ( + workers = 6 + duration = 10 * time.Second + ) + commands := [][]string{ + {"auth:token", "--no-warn"}, + {"auth:info", "id", "--refresh"}, + } + deadline := time.Now().Add(duration) + var ( + wg sync.WaitGroup + mu sync.Mutex + runs int + errs []string + ) + for i := range workers { + wg.Go(func() { + for time.Now().Before(deadline) { + _, stderr, err := f.RunCombinedOutput(commands[i%len(commands)]...) + mu.Lock() + runs++ + if err != nil { + errs = append(errs, stderr) + } + mu.Unlock() + } + }) + } + wg.Wait() + + assert.Empty(t, errs) + assert.False(t, authServer.ReuseDetected(), "a rotated refresh token was reused") + expiries := int(duration/(3*time.Second)) + 1 + refreshes := authServer.RefreshRequests() + t.Logf("%d commands, %d refreshes", runs, refreshes) + assert.GreaterOrEqual(t, refreshes, 2) + assert.LessOrEqual(t, refreshes, expiries+1, "there must be one refresh per expiry") +} + +// TestAuthRefresh_KilledWhileRefreshing checks that a process killed while holding the lock does not block others. +func TestAuthRefresh_KilledWhileRefreshing(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("uses SIGKILL") + } + authServer := mockapi.NewAuthServer(t) + defer authServer.Close() + authServer.AddRefreshToken("initial-refresh-token") + authServer.SetRefreshDelay(time.Minute) + apiServer := httptest.NewServer(mockapi.NewHandler(t)) + defer apiServer.Close() + + f := newCommandFactory(t, apiServer.URL, authServer.URL) + f.extraEnv = append(f.extraEnv, EnvPrefix+"TOKEN=") + writeOAuthSession(t, f.home, "default", map[string]any{ + "accessToken": "expired-token", + "expires": time.Now().Add(-time.Hour).Unix(), + "refreshToken": "initial-refresh-token", + }) + + cmd := f.buildCommand("auth:token", "--no-warn") + require.NoError(t, cmd.Start()) + require.Eventually(t, func() bool { return authServer.RefreshRequests() == 1 }, 20*time.Second, 50*time.Millisecond) + require.NoError(t, cmd.Process.Kill()) + _ = cmd.Wait() + + authServer.SetRefreshDelay(0) + start := time.Now() + assert.Equal(t, "access-token-1", f.Run("auth:token", "--no-warn")) + assert.Less(t, time.Since(start), 10*time.Second) + assert.False(t, authServer.ReuseDetected()) +} + +// TestAuthRefresh_TransientError checks that server errors during a refresh keep the session. +func TestAuthRefresh_TransientError(t *testing.T) { + authServer := mockapi.NewAuthServer(t) + defer authServer.Close() + authServer.AddRefreshToken("initial-refresh-token") + apiServer := httptest.NewServer(mockapi.NewHandler(t)) + defer apiServer.Close() + + f := newCommandFactory(t, apiServer.URL, authServer.URL) + f.extraEnv = append(f.extraEnv, EnvPrefix+"TOKEN=") + writeOAuthSession(t, f.home, "default", map[string]any{ + "accessToken": "expired-token", + "expires": time.Now().Add(-time.Hour).Unix(), + "refreshToken": "initial-refresh-token", + }) + + // A failure after the request was sent is not retried. + authServer.SetRefreshFailures(1) + _, stderr, err := f.RunCombinedOutput("auth:token", "--no-warn") + assertExitCode(t, 1, err) + assert.Contains(t, stderr, "failed to refresh the access token") + assert.Equal(t, 1, authServer.RefreshRequests()) + + // The session is kept, so the next refresh succeeds. + assert.Equal(t, "access-token-1", f.Run("auth:token", "--no-warn")) + assert.Equal(t, 2, authServer.RefreshRequests()) + assert.False(t, authServer.ReuseDetected()) +} + +// TestAuthRefresh_InvalidGrantClearsSessionFiles checks that an expired session's files are deleted. +func TestAuthRefresh_InvalidGrantClearsSessionFiles(t *testing.T) { + authServer := mockapi.NewAuthServer(t) + defer authServer.Close() + + f := newCommandFactory(t, "", authServer.URL) + f.extraEnv = append(f.extraEnv, EnvPrefix+"TOKEN=") + writeOAuthSession(t, f.home, "default", map[string]any{ + "accessToken": "expired-token", + "expires": time.Now().Add(-time.Hour).Unix(), + "refreshToken": "unknown-refresh-token", + }) + dir := filepath.Join(f.home, ".platform-test-cli") + certFile := filepath.Join(dir, ".session", "sess-cli-default", "ssh", "id_ed25519-cert.pub") + sshConfig := filepath.Join(dir, "ssh", "session.config") + require.NoError(t, os.WriteFile(certFile, []byte("cert"), 0o600)) + require.NoError(t, os.MkdirAll(filepath.Dir(sshConfig), 0o700)) + require.NoError(t, os.WriteFile(sshConfig, []byte("config"), 0o600)) + + _, stderr, err := f.RunCombinedOutput("auth:token", "--no-warn") + assertExitCode(t, 3, err) + assert.Contains(t, stderr, "logged out") + assert.NoFileExists(t, certFile) + assert.NoFileExists(t, sshConfig) +} + +// TestAuthStepUp checks the message for a step-up authentication challenge (RFC 9470) in a PHP command. +func TestAuthStepUp(t *testing.T) { + authServer := mockapi.NewAuthServer(t) + defer authServer.Close() + apiHandler := mockapi.NewHandler(t) + apiHandler.SetMyUser(&mockapi.User{ID: "u1"}) + apiHandler.RequireStepUp([]string{"mfa"}) + apiServer := httptest.NewServer(apiHandler) + defer apiServer.Close() + + f := newCommandFactory(t, apiServer.URL, authServer.URL) + f.extraEnv = append(f.extraEnv, EnvPrefix+"TOKEN=") + writeOAuthSession(t, f.home, "default", map[string]any{ + "accessToken": "valid-token", + "expires": time.Now().Add(time.Hour).Unix(), + }) + + _, stderr, err := f.RunCombinedOutput("auth:info", "id", "--refresh") + assertExitCode(t, 3, err) + assert.Contains(t, stderr, "Multi-factor authentication (MFA) is required.") + assert.Contains(t, stderr, "platform-test login --method mfa") +} + +// TestAuthInternal checks the JSON output of the hidden command used by the legacy CLI. +func TestAuthInternal(t *testing.T) { + authServer := mockapi.NewAuthServer(t) + defer authServer.Close() + + f := newCommandFactory(t, "", authServer.URL) + out := f.Run("auth:internal", "status") + assert.JSONEq(t, `{"logged_in": true, "session_ids": [], "has_stored_api_token": false}`, out) + + out = f.Run("auth:internal", "token") + assert.Contains(t, out, `"access_token":"access-token-1"`) + + f.extraEnv = append(f.extraEnv, EnvPrefix+"TOKEN=") + stdout, stderr, err := f.RunCombinedOutput("auth:internal", "token") + assertExitCode(t, 3, err) + assert.Empty(t, stdout) + assert.Equal(t, "{}", strings.TrimSpace(stderr)) +} + +// TestAuthInternalCommandsHidden checks that the internal legacy commands are hidden, and not run via abbreviations. +func TestAuthInternalCommandsHidden(t *testing.T) { + f := newCommandFactory(t, "", "") + + var list struct { + Commands []struct { + Name string `json:"name"` + Hidden bool `json:"hidden"` + } `json:"commands"` + } + require.NoError(t, json.Unmarshal([]byte(f.Run("list", "--all", "--format=json")), &list)) + hidden := map[string]bool{} + for _, c := range list.Commands { + hidden[c.Name] = c.Hidden + } + for _, name := range []string{"auth:export-sessions", "auth:post-login"} { + h, ok := hidden[name] + assert.True(t, ok && h, "%s must be listed as hidden", name) + } + + _, stderr, err := f.RunCombinedOutput("auth:ex") + assert.Error(t, err) + assert.Contains(t, stderr, `The command "auth:ex" does not exist.`) +} + +// TestErrorOutput_QuietOverride checks that --verbose and --debug override --quiet for errors. +func TestErrorOutput_QuietOverride(t *testing.T) { + f := newCommandFactory(t, "", "") + cases := []struct { + args []string + wantError bool + }{ + {[]string{"-q"}, false}, + {[]string{"-qv"}, true}, + {[]string{"-q", "--debug"}, true}, + } + for _, c := range cases { + _, stderr, err := f.RunCombinedOutput(append(append([]string{"auth:token"}, c.args...), "--unknown")...) + assert.Error(t, err) + if c.wantError { + assert.Contains(t, stderr, "Error:", "args: %v", c.args) + } else { + assert.NotContains(t, stderr, "Error:", "args: %v", c.args) + } + } +} diff --git a/integration-tests/auth_logout_test.go b/integration-tests/auth_logout_test.go index 884b4c3b3..e4313f582 100644 --- a/integration-tests/auth_logout_test.go +++ b/integration-tests/auth_logout_test.go @@ -89,15 +89,16 @@ func TestAuthLogout_Other(t *testing.T) { require.NoError(t, err, "stderr: %s", stderr) assert.Contains(t, stderr, "All other sessions have been deleted") - // "default" session file must still exist. - defaultSessFile := filepath.Join(f.home, ".platform-test-cli", ".session", "sess-default", "sess-default.json") - assert.FileExists(t, defaultSessFile) - - // The "other" session's tokens are revoked and its files deleted. - // An empty sess-other directory may be left behind. - assert.Contains(t, authServer.RevokedTokens(), "token-other") - assert.NotContains(t, authServer.RevokedTokens(), "token-default") + // The sessions were migrated from the legacy CLI's storage, which was then deleted. sessDir := filepath.Join(f.home, ".platform-test-cli", ".session") + assert.NoFileExists(t, filepath.Join(sessDir, "sess-default", "sess-default.json")) assert.NoFileExists(t, filepath.Join(sessDir, "sess-other", "sess-other.json")) assert.NoDirExists(t, filepath.Join(sessDir, "sess-cli-other")) + + // Only the "default" session remains, and only the "other" session's tokens are revoked. + authDir := filepath.Join(f.home, ".platform-test-cli", "auth") + assert.FileExists(t, filepath.Join(authDir, "default.json")) + assert.NoFileExists(t, filepath.Join(authDir, "other.json")) + assert.Contains(t, authServer.RevokedTokens(), "token-other") + assert.NotContains(t, authServer.RevokedTokens(), "token-default") } diff --git a/integration-tests/tests.go b/integration-tests/tests.go index 3db1ccb12..20eb73dfd 100644 --- a/integration-tests/tests.go +++ b/integration-tests/tests.go @@ -192,7 +192,26 @@ func (f *cmdFactory) fakeBrowser() { f.t.Skip("the fake browser is a shell script") } dir := f.t.TempDir() - require.NoError(f.t, os.WriteFile(filepath.Join(dir, "xdg-open"), []byte("#!/bin/sh\nexit 0\n"), 0o755)) + // The CLI uses "open" on macOS, and "xdg-open" on Linux. + for _, name := range []string{"open", "xdg-open"} { + require.NoError(f.t, os.WriteFile(filepath.Join(dir, name), []byte("#!/bin/sh\nexit 0\n"), 0o755)) + } + f.extraEnv = append(f.extraEnv, "DISPLAY=:0", "PATH="+dir+string(os.PathListSeparator)+os.Getenv("PATH")) +} + +// fakeBrowserThatLogsIn makes the CLI detect a display and a browser, which completes the login flow with the mock +// auth server by following its redirects. +func (f *cmdFactory) fakeBrowserThatLogsIn() { + f.t.Helper() + if runtime.GOOS == "windows" { + f.t.Skip("the fake browser is a shell script") + } + dir := f.t.TempDir() + script := "#!/bin/sh\nexec curl -fsSL -o /dev/null \"$1\"\n" + // The CLI uses "open" on macOS, and "xdg-open" on Linux. + for _, name := range []string{"open", "xdg-open"} { + require.NoError(f.t, os.WriteFile(filepath.Join(dir, name), []byte(script), 0o755)) + } f.extraEnv = append(f.extraEnv, "DISPLAY=:0", "PATH="+dir+string(os.PathListSeparator)+os.Getenv("PATH")) } diff --git a/internal/auth/client.go b/internal/auth/client.go deleted file mode 100644 index 8d1f77129..000000000 --- a/internal/auth/client.go +++ /dev/null @@ -1,54 +0,0 @@ -package auth - -import ( - "context" - "fmt" - "net/http" - - "golang.org/x/oauth2" - - "github.com/upsun/cli/internal/legacy" -) - -type LegacyCLIClient struct { - HTTPClient *http.Client - tokenSource oauth2.TokenSource -} - -func (c *LegacyCLIClient) EnsureAuthenticated(_ context.Context) error { - _, err := c.tokenSource.Token() - return err -} - -// NewLegacyCLIClient creates an HTTP client authenticated through the legacy CLI. -// The wrapper argument must be a dedicated wrapper, not used by other callers. -func NewLegacyCLIClient(ctx context.Context, wrapper *legacy.CLIWrapper) (*LegacyCLIClient, error) { - ts, err := NewLegacyCLITokenSource(ctx, wrapper) - if err != nil { - return nil, fmt.Errorf("oauth2: create token source: %w", err) - } - - refresher, ok := ts.(refresher) - if !ok { - return nil, fmt.Errorf("token source does not implement refresher") - } - baseRT := http.DefaultTransport - if rt, ok := TransportFromContext(ctx); ok && rt != nil { - baseRT = rt - } - - httpClient := &http.Client{ - Transport: &Transport{ - refresher: refresher, - base: &oauth2.Transport{ - Source: ts, - Base: baseRT, - }, - }, - } - - return &LegacyCLIClient{ - HTTPClient: httpClient, - tokenSource: ts, - }, nil -} diff --git a/internal/auth/errors.go b/internal/auth/errors.go new file mode 100644 index 000000000..e9b24e960 --- /dev/null +++ b/internal/auth/errors.go @@ -0,0 +1,84 @@ +package auth + +import ( + "encoding/json" + "io" + "net/http" + "strings" +) + +// ExitCodeLoginRequired is the exit code used when login is required. +const ExitCodeLoginRequired = 3 + +// LoginRequiredError means the user must log in (again). +type LoginRequiredError struct { + // Notice is shown before the login prompt, e.g. "Your session has expired. You have been logged out." + Notice string `json:"notice,omitempty"` + + // AuthMethods and MaxAge come from a step-up authentication challenge (RFC 9470). + AuthMethods []string `json:"amr,omitempty"` + MaxAge *int `json:"max_age,omitempty"` + + HasAPIToken bool `json:"has_api_token,omitempty"` +} + +func (e *LoginRequiredError) Error() string { + return e.Message() +} + +// Message returns a short description, matching the legacy CLI's LoginRequiredEvent::getMessage(). +func (e *LoginRequiredError) Message() string { + msg := "Authentication is required." + if len(e.AuthMethods) > 0 || e.MaxAge != nil { + msg = "Re-authentication is required." + } + switch { + case len(e.AuthMethods) == 1 && e.AuthMethods[0] == "mfa": + msg = "Multi-factor authentication (MFA) is required." + case len(e.AuthMethods) == 1 && strings.HasPrefix(e.AuthMethods[0], "sso:"): + msg = "Single sign-on (SSO) is required." + case len(e.AuthMethods) != 1 && e.MaxAge != nil: + msg = "More recent authentication is required." + } + return msg +} + +// loginRequiredAfterRefreshError converts a failed refresh into a login-required error, with the legacy CLI's notices. +func loginRequiredAfterRefreshError(oerr *OAuthError) *LoginRequiredError { + switch { + case strings.Contains(oerr.Description, "SSO session has expired"): + return &LoginRequiredError{Notice: "Your SSO session has expired. You have been logged out."} + case strings.Contains(oerr.Description, "API token"): + return &LoginRequiredError{Notice: "The API token is invalid.", HasAPIToken: true} + default: + return &LoginRequiredError{Notice: "Your session has expired. You have been logged out."} + } +} + +// isInvalidAPITokenError checks if a token endpoint error means the API token was rejected. +func isInvalidAPITokenError(oerr *OAuthError) bool { + if oerr.StatusCode != http.StatusBadRequest && oerr.StatusCode != http.StatusUnauthorized { + return false + } + return oerr.Code == "invalid_grant" || oerr.Code == "request_unauthorized" +} + +// IsStepUpChallenge checks for a step-up authentication response (RFC 9470). +func IsStepUpChallenge(resp *http.Response) bool { + if resp.StatusCode != http.StatusUnauthorized { + return false + } + h := strings.Join(resp.Header.Values("WWW-Authenticate"), "\n") + return strings.Contains(strings.ToLower(h), "bearer") && strings.Contains(h, "insufficient_user_authentication") +} + +// StepUpError reads the required authentication methods and max age from a step-up response body. +func StepUpError(resp *http.Response, hasAPIToken bool) *LoginRequiredError { + var body struct { + AMR []string `json:"amr"` + MaxAge *int `json:"max_age"` + } + b, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20)) + _ = json.Unmarshal(b, &body) + return &LoginRequiredError{AuthMethods: body.AMR, MaxAge: body.MaxAge, HasAPIToken: hasAPIToken} +} diff --git a/internal/auth/legacy.go b/internal/auth/legacy.go deleted file mode 100644 index e07ee662f..000000000 --- a/internal/auth/legacy.go +++ /dev/null @@ -1,108 +0,0 @@ -package auth - -import ( - "bytes" - "context" - "fmt" - "io" - "sync" - - "golang.org/x/oauth2" - - "github.com/upsun/cli/internal/legacy" -) - -type legacyCLITokenSource struct { - ctx context.Context - cached *oauth2.Token - wrapper *legacy.CLIWrapper - mu sync.Mutex -} - -func (ts *legacyCLITokenSource) unsafeGetLegacyCLIToken() (*oauth2.Token, error) { - bt := bytes.NewBuffer(nil) - ts.wrapper.Stdout = bt - if err := ts.wrapper.Exec(ts.ctx, "auth:token", "-W"); err != nil { - return nil, fmt.Errorf("cannot retrieve token: %w", err) - } - - expiry, err := unsafeGetJWTExpiry(bt.String()) - - if err != nil { - return nil, fmt.Errorf("cannot parse token: %w", err) - } - - return &oauth2.Token{ - AccessToken: bt.String(), - TokenType: "Bearer", - Expiry: expiry, - }, nil -} - -func (ts *legacyCLITokenSource) refreshToken() error { - ts.mu.Lock() - defer ts.mu.Unlock() - - return ts.unsafeRefreshToken() -} - -func (ts *legacyCLITokenSource) unsafeRefreshToken() error { - ts.cached = nil - ts.wrapper.Stdout = io.Discard - if err := ts.wrapper.Exec(ts.ctx, "auth:info", "--refresh"); err != nil { - return fmt.Errorf("cannot refresh token: %w", err) - } - - return nil -} - -func (ts *legacyCLITokenSource) invalidateToken() error { - ts.mu.Lock() - defer ts.mu.Unlock() - - return ts.unsafeInvalidateToken() -} - -func (ts *legacyCLITokenSource) unsafeInvalidateToken() error { - if ts.cached != nil { - ts.cached.AccessToken = "" - } - - return nil -} - -func (ts *legacyCLITokenSource) Token() (*oauth2.Token, error) { - ts.mu.Lock() - defer ts.mu.Unlock() - - if ts.cached == nil { - tok, err := ts.unsafeGetLegacyCLIToken() - if err != nil { - return nil, err - } - ts.cached = tok - } - - if ts.cached != nil && ts.cached.Valid() { - return ts.cached, nil - } - - if err := ts.unsafeRefreshToken(); err != nil { - return nil, err - } - - tok, err := ts.unsafeGetLegacyCLIToken() - if err != nil { - return nil, err - } - - ts.cached = tok - return ts.cached, nil -} - -func NewLegacyCLITokenSource(ctx context.Context, wrapper *legacy.CLIWrapper) (oauth2.TokenSource, error) { - return &legacyCLITokenSource{ - ctx: ctx, - wrapper: wrapper, - }, nil -} diff --git a/internal/auth/manager.go b/internal/auth/manager.go new file mode 100644 index 000000000..4881afe58 --- /dev/null +++ b/internal/auth/manager.go @@ -0,0 +1,520 @@ +package auth + +import ( + "context" + "crypto/sha256" + "crypto/tls" + "encoding/hex" + "errors" + "fmt" + "io" + "net/http" + "os" + "path/filepath" + "strings" + "sync" + "time" + + "github.com/gofrs/flock" + + "github.com/upsun/cli/internal/auth/store" + "github.com/upsun/cli/internal/config" +) + +// expiryMargin is how long before its expiry a token counts as expired, so that requests rarely start with a token +// that expires in flight. +const expiryMargin = 2 * time.Minute + +// defaultLockWait bounds how long to wait for another process's refresh. +const defaultLockWait = 60 * time.Second + +// Token is an access token and its expiry time. +type Token struct { + AccessToken string `json:"access_token"` + Expires int64 `json:"expires,omitempty"` // A Unix timestamp, or 0 if unknown. +} + +// Status describes the authentication state, without making any network requests. +type Status struct { + LoggedIn bool `json:"logged_in"` + SessionIDs []string `json:"session_ids"` + HasStoredAPIToken bool `json:"has_stored_api_token"` +} + +// Manager handles tokens and sessions. It is the only component that reads or writes credentials. +type Manager struct { + Settings *config.Auth + Store *store.Store + OAuth *OAuthClient + + // Migrator imports credentials from the legacy CLI's storage, once. It may be nil. + Migrator *Migrator + + // Stderr receives warnings. + Stderr io.Writer + + // OnLoggedOut is called after a refresh fails and the session is deleted, e.g. to delete other session files. + // It may be nil. + OnLoggedOut func(id string) error + + // LockWait bounds how long to wait for a lock. It defaults to 60s. + LockWait time.Duration + + migrateOnce sync.Once + migrateErr error + warnedLocks sync.Once +} + +// NewManager creates a Manager from the CLI config. +func NewManager(cnf *config.Config, stderr io.Writer) (*Manager, error) { + settings, err := cnf.Auth() + if err != nil { + return nil, err + } + dir, err := cnf.WritableUserDir() //nolint:staticcheck // credentials belong in the user dir, not a cache + if err != nil { + return nil, err + } + httpClient := NewHTTPClient(cnf, settings) + return &Manager{ + Settings: settings, + Store: &store.Store{ + Dir: filepath.Join(dir, "auth"), + Service: cnf.Application.Slug + "-cli-auth", + UseKeychain: !settings.DisableCredentialHelpers && store.KeychainSupported(), + Stderr: stderr, + }, + OAuth: &OAuthClient{ + HTTPClient: httpClient, + TokenURL: settings.TokenURL, + RevokeURL: settings.RevokeURL, + ClientID: settings.ClientID, + }, + Stderr: stderr, + }, nil +} + +// NewHTTPClient returns an HTTP client for the API and auth servers, without authentication. +// +// It uses Go's defaults for proxies (HTTPS_PROXY, NO_PROXY) and the system trust store. +func NewHTTPClient(cnf *config.Config, settings *config.Auth) *http.Client { + t := http.DefaultTransport.(*http.Transport).Clone() //nolint:errcheck // the default is a *http.Transport + if settings.SkipSSL { + t.TLSClientConfig = &tls.Config{InsecureSkipVerify: true} //nolint:gosec // the user disabled verification + } + return &http.Client{Transport: &userAgentTransport{base: t, userAgent: cnf.UserAgent()}} +} + +type userAgentTransport struct { + base http.RoundTripper + userAgent string +} + +func (t *userAgentTransport) RoundTrip(req *http.Request) (*http.Response, error) { + if req.Header.Get("User-Agent") == "" { + req = req.Clone(req.Context()) + req.Header.Set("User-Agent", t.userAgent) + } + return t.base.RoundTrip(req) +} + +// APITokenSessionID returns the session ID used for the OAuth 2.0 tokens obtained from an API token. +// +// This keeps them separate from the user's session, as in the legacy CLI. +func APITokenSessionID(apiToken string) string { + sum := sha256.Sum256([]byte(apiToken)) + return "api-token-" + hex.EncodeToString(sum[:])[:32] +} + +// ConfiguredAPIToken returns an API token set via config (api.token or api.token_file), if any. +func (m *Manager) ConfiguredAPIToken() (string, error) { + if m.Settings.Token != "" { + return m.Settings.Token, nil + } + if m.Settings.TokenFile != "" { + b, err := os.ReadFile(m.Settings.TokenFile) + if err != nil { + return "", fmt.Errorf("failed to read file: %s", m.Settings.TokenFile) + } + return strings.TrimSpace(string(b)), nil + } + return "", nil +} + +// HasConfiguredToken reports whether an API token or access token is set via config. +func (m *Manager) HasConfiguredToken() (bool, error) { + if m.Settings.AccessToken != "" { + return true, nil + } + t, err := m.ConfiguredAPIToken() + return t != "", err +} + +// apiToken returns the API token to use, in order of precedence: a stored one, then one set via config. +func (m *Manager) apiToken(ctx context.Context) (string, error) { + e, err := m.load(ctx, m.Settings.SessionID) + if err != nil { + return "", err + } + if e != nil && e.APIToken != "" { + return e.APIToken, nil + } + return m.ConfiguredAPIToken() +} + +// Token returns a valid access token, refreshing it if needed. +// +// If rejected is set (an access token the API rejected), and it is still the stored token, the token is refreshed. +func (m *Manager) Token(ctx context.Context, rejected string) (*Token, error) { + apiToken, err := m.apiToken(ctx) + if err != nil { + return nil, err + } + if apiToken != "" { + return m.refresh(ctx, APITokenSessionID(apiToken), rejected, apiToken) + } + if m.Settings.AccessToken != "" { + t := &Token{AccessToken: m.Settings.AccessToken} + if exp, err := unsafeGetJWTExpiry(t.AccessToken); err == nil { + t.Expires = exp.Unix() + } + return t, nil + } + return m.refresh(ctx, m.Settings.SessionID, rejected, "") +} + +// Status returns the authentication state. +func (m *Manager) Status(ctx context.Context) (*Status, error) { + s := &Status{SessionIDs: []string{}} + ids, err := m.SessionIDs(ctx) + if err != nil { + return nil, err + } + s.SessionIDs = ids + e, err := m.load(ctx, m.Settings.SessionID) + if err != nil { + return nil, err + } + s.HasStoredAPIToken = e != nil && e.APIToken != "" + hasConfigured, err := m.HasConfiguredToken() + if err != nil { + return nil, err + } + s.LoggedIn = hasConfigured || (e != nil && (e.AccessToken != "" || e.RefreshToken != "" || e.APIToken != "")) + return s, nil +} + +// SessionIDs lists the user's sessions, excluding the sessions for API tokens. +func (m *Manager) SessionIDs(ctx context.Context) ([]string, error) { + if err := m.migrate(ctx); err != nil { + return nil, err + } + ids, err := m.Store.List() + if err != nil { + return nil, err + } + out := []string{} + for _, id := range ids { + if !strings.HasPrefix(id, "api-token-") { + out = append(out, id) + } + } + return out, nil +} + +// Load returns the stored entry for a session, or nil. +func (m *Manager) Load(ctx context.Context, sessionID string) (*store.Entry, error) { + return m.load(ctx, sessionID) +} + +func (m *Manager) load(ctx context.Context, id string) (*store.Entry, error) { + if err := m.migrate(ctx); err != nil { + return nil, err + } + return m.Store.Load(id) +} + +// Save saves the entry for a session, under the session's lock. +func (m *Manager) Save(ctx context.Context, id string, e *store.Entry) error { + if err := m.migrate(ctx); err != nil { + return err + } + unlock, err := m.lock(ctx, id) + if err != nil { + return err + } + defer unlock() + return m.Store.Save(id, e) +} + +// refresh returns a valid token for the session, refreshing it under the session's lock. +// +// Refresh tokens rotate on every refresh, and sending a rotated token again revokes the whole login. So the refresh +// token is always read from the store while the lock is held. +func (m *Manager) refresh(ctx context.Context, id, rejected, apiToken string) (*Token, error) { + if err := m.migrate(ctx); err != nil { + return nil, err + } + unlock, err := m.lock(ctx, id) + if err != nil { + return nil, err + } + defer unlock() + + e, err := m.Store.Load(id) + if err != nil { + return nil, err + } + if e != nil && usable(e, rejected) { + return tokenFromEntry(e), nil + } + + if e != nil && e.RefreshToken != "" { + newEntry, err := m.OAuth.Refresh(ctx, e.RefreshToken) + var oerr *OAuthError + switch { + case err == nil: + if newEntry.RefreshToken == "" { + newEntry.RefreshToken = e.RefreshToken + } + // Save the new tokens before doing anything else. + if err := m.Store.Save(id, newEntry); err != nil { + return nil, err + } + return tokenFromEntry(newEntry), nil + case errors.As(err, &oerr) && oerr.Code == "invalid_request": + // Another process sent the same refresh token without the lock. The session is kept. + return m.afterConcurrentRefresh(ctx, id, e) + case errors.As(err, &oerr) && oerr.Code == "invalid_grant": + if err := m.Store.Delete(id); err != nil { + return nil, err + } + if apiToken == "" { + if err := m.loggedOut(id); err != nil { + return nil, err + } + return nil, loginRequiredAfterRefreshError(oerr) + } + // Exchange the API token again below. + default: + return nil, fmt.Errorf("failed to refresh the access token: %w", err) + } + } + + if apiToken != "" { + newEntry, err := m.OAuth.ExchangeAPIToken(ctx, apiToken) + if err != nil { + var oerr *OAuthError + if errors.As(err, &oerr) && isInvalidAPITokenError(oerr) { + return nil, &LoginRequiredError{Notice: "The API token is invalid.", HasAPIToken: true} + } + return nil, fmt.Errorf("failed to exchange the API token: %w", err) + } + if err := m.Store.Save(id, newEntry); err != nil { + return nil, err + } + return tokenFromEntry(newEntry), nil + } + + if e != nil && e.AccessToken != "" { + // An expired or rejected token, with no way to refresh it. + if err := m.Store.Delete(id); err != nil { + return nil, err + } + if err := m.loggedOut(id); err != nil { + return nil, err + } + return nil, &LoginRequiredError{Notice: "Your session has expired. You have been logged out."} + } + return nil, &LoginRequiredError{} +} + +func (m *Manager) loggedOut(id string) error { + if m.OnLoggedOut == nil { + return nil + } + return m.OnLoggedOut(id) +} + +// afterConcurrentRefresh waits briefly and uses the stored tokens if another process saved new ones. +func (m *Manager) afterConcurrentRefresh(ctx context.Context, id string, old *store.Entry) (*Token, error) { + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-time.After(time.Second): + } + e, err := m.Store.Load(id) + if err != nil { + return nil, err + } + if e != nil && e.RefreshToken != old.RefreshToken && e.AccessToken != "" { + return tokenFromEntry(e), nil + } + return nil, errors.New("the access token is being refreshed by another process: please try again") +} + +// usable reports whether an entry's access token can be used without refreshing. +func usable(e *store.Entry, rejected string) bool { + if e.AccessToken == "" || e.AccessToken == rejected { + return false + } + return e.Expires == 0 || time.Unix(e.Expires, 0).After(time.Now().Add(expiryMargin)) +} + +func tokenFromEntry(e *store.Entry) *Token { + return &Token{AccessToken: e.AccessToken, Expires: e.Expires} +} + +// lock takes the session's lock, bounded by the context and LockWait. +func (m *Manager) lock(ctx context.Context, id string) (unlock func(), err error) { + mu := sessionMutex(id) + mu.Lock() + if m.Settings.DisableLocks { + m.warnedLocks.Do(func() { + if m.Stderr != nil { + fmt.Fprintln(m.Stderr, "Warning: locks are disabled (api.disable_locks). "+ + "Concurrent commands can end the session.") + } + }) + return mu.Unlock, nil + } + fileUnlock, err := m.fileLock(ctx, m.Store.LockPath(id)) + if err != nil { + mu.Unlock() + return nil, err + } + return func() { + fileUnlock() + mu.Unlock() + }, nil +} + +// fileLock takes an exclusive OS-level lock. The OS releases it if the process dies. +func (m *Manager) fileLock(ctx context.Context, path string) (unlock func(), err error) { + if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil { + return nil, err + } + wait := m.LockWait + if wait == 0 { + wait = defaultLockWait + } + ctx, cancel := context.WithTimeout(ctx, wait) + defer cancel() + fl := flock.New(path, flock.SetPermissions(0o600)) + ok, err := fl.TryLockContext(ctx, 50*time.Millisecond) + if err != nil || !ok { + if errors.Is(err, context.DeadlineExceeded) { + return nil, errors.New("timed out waiting for another process to refresh the access token: please try again") + } + return nil, fmt.Errorf("failed to lock %s: %w", path, err) + } + return func() { _ = fl.Unlock() }, nil +} + +var ( + sessionMutexes = map[string]*sync.Mutex{} + sessionMutexesMu sync.Mutex +) + +// sessionMutex returns a mutex per session ID, so goroutines in one process share one refresh. +func sessionMutex(id string) *sync.Mutex { + sessionMutexesMu.Lock() + defer sessionMutexesMu.Unlock() + mu, ok := sessionMutexes[id] + if !ok { + mu = &sync.Mutex{} + sessionMutexes[id] = mu + } + return mu +} + +// Logout revokes and deletes a session, including the session for its saved API token. +// If the session is the current one, the session for a configured API token is also logged out. +func (m *Manager) Logout(ctx context.Context, id string) error { + e, err := m.load(ctx, id) + if err != nil { + return err + } + var apiTokenSessions []string + if e != nil && e.APIToken != "" { + apiTokenSessions = append(apiTokenSessions, APITokenSessionID(e.APIToken)) + } + if id == m.Settings.SessionID { + if t, _ := m.ConfiguredAPIToken(); t != "" { + apiTokenSessions = append(apiTokenSessions, APITokenSessionID(t)) + } + } + var errs []error + for _, sid := range apiTokenSessions { + errs = append(errs, m.logoutOne(ctx, sid)) + } + errs = append(errs, m.logoutOne(ctx, id)) + return errors.Join(errs...) +} + +// LogoutToReplace logs out of a session before new credentials are saved to it, with apiToken if it is an API token +// login. If the keychain cannot be used, the session files are forgotten instead, so that new ones can be saved. +func (m *Manager) LogoutToReplace(ctx context.Context, id, apiToken string) error { + err := m.Logout(ctx, id) + var kerr *store.KeychainError + if err == nil || !errors.As(err, &kerr) { + return err + } + if m.Stderr != nil { + fmt.Fprintf(m.Stderr, "Warning: %s\n", err) + } + ids := []string{id} + if apiToken != "" { + ids = append(ids, APITokenSessionID(apiToken)) + } + for _, sid := range ids { + if err := m.Store.Forget(sid); err != nil { + return err + } + } + return nil +} + +func (m *Manager) logoutOne(ctx context.Context, id string) error { + unlock, err := m.lock(ctx, id) + if err != nil { + return err + } + defer unlock() + e, err := m.Store.Load(id) + if err != nil { + return err + } + if e != nil { + for _, r := range []struct{ token, hint string }{ + {e.RefreshToken, "refresh_token"}, + {e.AccessToken, "access_token"}, + } { + if r.token == "" { + continue + } + if err := m.OAuth.Revoke(ctx, r.token, r.hint); err != nil && m.Stderr != nil { + fmt.Fprintf(m.Stderr, "Warning: failed to revoke the %s: %s\n", strings.ReplaceAll(r.hint, "_", " "), err) + } + } + } + return m.Store.Delete(id) +} + +// DeleteAll deletes every session's credentials, without revoking them. +func (m *Manager) DeleteAll(ctx context.Context) error { + if err := m.migrate(ctx); err != nil { + return err + } + return m.Store.DeleteAll() +} + +func (m *Manager) migrate(ctx context.Context) error { + if m.Migrator == nil { + return nil + } + m.migrateOnce.Do(func() { + m.migrateErr = m.Migrator.Run(ctx, m) + }) + return m.migrateErr +} diff --git a/internal/auth/manager_test.go b/internal/auth/manager_test.go new file mode 100644 index 000000000..4c0385d75 --- /dev/null +++ b/internal/auth/manager_test.go @@ -0,0 +1,581 @@ +package auth + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/zalando/go-keyring" + + "github.com/upsun/cli/internal/auth/store" + "github.com/upsun/cli/internal/config" +) + +// testAuthServer is an OAuth 2.0 token endpoint that rotates refresh tokens and detects reuse. +type testAuthServer struct { + *httptest.Server + mu sync.Mutex + refreshes int + issued int + valid map[string]bool // refresh token → unused + reused bool + revoked []string + failNext []int // status codes to return for the next refresh requests + errorNext []string // OAuth error codes to return for the next refresh requests + dropNext int // the number of refresh requests to accept, but whose responses are dropped + lifetime int64 + delay time.Duration +} + +func newTestAuthServer(t *testing.T) *testAuthServer { + s := &testAuthServer{valid: map[string]bool{}, lifetime: 3600} + s.Server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.NoError(t, r.ParseForm()) + if r.URL.Path == "/revoke" { + s.mu.Lock() + s.revoked = append(s.revoked, r.Form.Get("token")) + s.mu.Unlock() + return + } + s.mu.Lock() + delay := s.delay + s.mu.Unlock() + time.Sleep(delay) + s.mu.Lock() + defer s.mu.Unlock() + writeErr := func(status int, code string) { + w.WriteHeader(status) + _ = json.NewEncoder(w).Encode(map[string]string{"error": code, "error_description": "test error " + code}) + } + switch r.Form.Get("grant_type") { + case "refresh_token": + s.refreshes++ + if len(s.failNext) > 0 { + status := s.failNext[0] + s.failNext = s.failNext[1:] + writeErr(status, "server_error") + return + } + if len(s.errorNext) > 0 { + code := s.errorNext[0] + s.errorNext = s.errorNext[1:] + writeErr(http.StatusBadRequest, code) + return + } + rt := r.Form.Get("refresh_token") + unused, known := s.valid[rt] + if !known || !unused { + if known { + s.reused = true + } + writeErr(http.StatusBadRequest, "invalid_grant") + return + } + s.valid[rt] = false + if s.dropNext > 0 { + s.dropNext-- + panic(http.ErrAbortHandler) + } + case "api_token": + if r.Form.Get("api_token") != "good-api-token" { + writeErr(http.StatusBadRequest, "request_unauthorized") + return + } + default: + writeErr(http.StatusBadRequest, "unsupported_grant_type") + return + } + s.issued++ + newRT := fmt.Sprintf("rt-%d", s.issued) + s.valid[newRT] = true + _ = json.NewEncoder(w).Encode(map[string]any{ + "access_token": fmt.Sprintf("at-%d", s.issued), + "refresh_token": newRT, + "token_type": "bearer", + "expires_in": s.lifetime, + }) + })) + t.Cleanup(s.Close) + return s +} + +func (s *testAuthServer) addRefreshToken(rt string) { + s.mu.Lock() + defer s.mu.Unlock() + s.valid[rt] = true +} + +func newTestManager(srv *testAuthServer, dir string, settings *config.Auth) *Manager { + s := *settings + if s.SessionID == "" { + s.SessionID = "default" + } + return &Manager{ + Settings: &s, + Store: &store.Store{Dir: dir}, + OAuth: &OAuthClient{ + HTTPClient: srv.Client(), + TokenURL: srv.URL + "/token", + RevokeURL: srv.URL + "/revoke", + ClientID: "test", + retryDelay: time.Millisecond, + }, + } +} + +func TestManager_Token(t *testing.T) { + future := time.Now().Add(time.Hour).Unix() + past := time.Now().Add(-time.Hour).Unix() + cases := []struct { + name string + entry *store.Entry + settings config.Auth + rejected string + setup func(s *testAuthServer) + wantToken string + wantRefreshes int + wantLogin string // The expected notice of a login-required error. + wantErr string + wantDeleted bool + }{ + { + name: "valid token", + entry: &store.Entry{AccessToken: "stored", RefreshToken: "rt-0", Expires: future}, + wantToken: "stored", + }, + { + name: "expired token", + entry: &store.Entry{AccessToken: "stored", RefreshToken: "rt-0", Expires: past}, + wantToken: "at-1", + wantRefreshes: 1, + }, + { + name: "inside the expiry margin", + entry: &store.Entry{ + AccessToken: "stored", RefreshToken: "rt-0", Expires: time.Now().Add(time.Minute).Unix(), + }, + wantToken: "at-1", + wantRefreshes: 1, + }, + { + name: "rejected token", + entry: &store.Entry{AccessToken: "stored", RefreshToken: "rt-0", Expires: future}, + rejected: "stored", + wantToken: "at-1", + wantRefreshes: 1, + }, + { + name: "rejected token already replaced", + entry: &store.Entry{AccessToken: "newer", RefreshToken: "rt-0", Expires: future}, + rejected: "older", + wantToken: "newer", + }, + { + name: "not logged in", + wantLogin: "", + }, + { + name: "invalid grant", + entry: &store.Entry{AccessToken: "stored", RefreshToken: "unknown", Expires: past}, + wantRefreshes: 1, + wantLogin: "Your session has expired. You have been logged out.", + wantDeleted: true, + }, + { + name: "expired without refresh token", + entry: &store.Entry{AccessToken: "stored", Expires: past}, + wantLogin: "Your session has expired. You have been logged out.", + wantDeleted: true, + }, + { + name: "5xx keeps the session", + entry: &store.Entry{AccessToken: "stored", RefreshToken: "rt-0", Expires: past}, + setup: func(s *testAuthServer) { s.failNext = []int{503} }, + wantRefreshes: 1, + wantErr: "failed to refresh the access token", + }, + { + name: "a lost response is not retried", + entry: &store.Entry{AccessToken: "stored", RefreshToken: "rt-0", Expires: past}, + setup: func(s *testAuthServer) { s.dropNext = 1 }, + wantRefreshes: 1, + wantErr: "failed to refresh the access token", + }, + { + name: "concurrent use keeps the session", + entry: &store.Entry{AccessToken: "stored", RefreshToken: "rt-0", Expires: past}, + setup: func(s *testAuthServer) { s.errorNext = []string{"invalid_request"} }, + wantRefreshes: 1, + wantErr: "being refreshed by another process", + }, + { + name: "access token from config", + settings: config.Auth{AccessToken: "raw"}, + wantToken: "raw", + }, + { + name: "API token from config", + settings: config.Auth{Token: "good-api-token"}, + wantToken: "at-1", + }, + { + name: "invalid API token", + settings: config.Auth{Token: "bad-api-token"}, + wantLogin: "The API token is invalid.", + }, + { + name: "stored API token takes precedence", + entry: &store.Entry{APIToken: "good-api-token"}, + settings: config.Auth{Token: "bad-api-token"}, + wantToken: "at-1", + }, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + srv := newTestAuthServer(t) + srv.addRefreshToken("rt-0") + if c.setup != nil { + c.setup(srv) + } + m := newTestManager(srv, t.TempDir(), &c.settings) + if c.entry != nil { + require.NoError(t, m.Store.Save("default", c.entry)) + } + var loggedOut []string + m.OnLoggedOut = func(id string) error { + loggedOut = append(loggedOut, id) + return nil + } + + tok, err := m.Token(context.Background(), c.rejected) + switch { + case c.wantErr != "": + assert.ErrorContains(t, err, c.wantErr) + case c.wantToken == "": + lerr, ok := AsLoginRequired(err) + require.True(t, ok, "expected a login-required error, got: %v", err) + assert.Equal(t, c.wantLogin, lerr.Notice) + default: + require.NoError(t, err) + assert.Equal(t, c.wantToken, tok.AccessToken) + } + srv.mu.Lock() + assert.Equal(t, c.wantRefreshes, srv.refreshes) + assert.False(t, srv.reused) + srv.mu.Unlock() + + stored, err := m.Store.Load("default") + require.NoError(t, err) + if c.wantDeleted { + assert.Nil(t, stored) + assert.Equal(t, []string{"default"}, loggedOut) + } else if c.entry != nil && c.entry.RefreshToken != "" && c.wantErr != "" { + assert.Equal(t, c.entry.RefreshToken, stored.RefreshToken, "the session must be kept") + } + if !c.wantDeleted { + assert.Empty(t, loggedOut) + } + }) + } +} + +// TestManager_ConcurrentRefresh simulates several processes, each with its own Manager, refreshing one session. +func TestManager_ConcurrentRefresh(t *testing.T) { + srv := newTestAuthServer(t) + srv.addRefreshToken("rt-0") + srv.delay = 50 * time.Millisecond + dir := t.TempDir() + require.NoError(t, newTestManager(srv, dir, &config.Auth{}).Store.Save("default", &store.Entry{ + AccessToken: "expired", RefreshToken: "rt-0", Expires: time.Now().Add(-time.Hour).Unix(), + })) + + const n = 10 + var wg sync.WaitGroup + tokens := make([]string, n) + errs := make([]error, n) + for i := range n { + wg.Go(func() { + m := newTestManager(srv, dir, &config.Auth{}) + tok, err := m.Token(context.Background(), "") + errs[i] = err + if tok != nil { + tokens[i] = tok.AccessToken + } + }) + } + wg.Wait() + for i := range n { + require.NoError(t, errs[i]) + assert.Equal(t, "at-1", tokens[i]) + } + assert.Equal(t, 1, srv.refreshes) + assert.False(t, srv.reused) +} + +func TestManager_LockTimeout(t *testing.T) { + srv := newTestAuthServer(t) + dir := t.TempDir() + m := newTestManager(srv, dir, &config.Auth{}) + m.LockWait = 100 * time.Millisecond + require.NoError(t, m.Store.Save("default", &store.Entry{ + AccessToken: "expired", RefreshToken: "rt-0", Expires: time.Now().Add(-time.Hour).Unix(), + })) + + // Another "process" holds the lock. + other := newTestManager(srv, dir, &config.Auth{}) + unlock, err := other.fileLock(context.Background(), m.Store.LockPath("default")) + require.NoError(t, err) + defer unlock() + + _, err = m.Token(context.Background(), "") + assert.ErrorContains(t, err, "timed out waiting") + assert.Equal(t, 0, srv.refreshes, "the token must never be refreshed without the lock") +} + +func TestManager_LogoutAndStatus(t *testing.T) { + srv := newTestAuthServer(t) + m := newTestManager(srv, t.TempDir(), &config.Auth{}) + ctx := context.Background() + + s, err := m.Status(ctx) + require.NoError(t, err) + assert.Equal(t, &Status{SessionIDs: []string{}}, s) + + require.NoError(t, m.Store.Save("default", &store.Entry{ + AccessToken: "at", RefreshToken: "rt", APIToken: "good-api-token", + })) + require.NoError(t, m.Store.Save(APITokenSessionID("good-api-token"), &store.Entry{AccessToken: "api-at"})) + require.NoError(t, m.Store.Save("other", &store.Entry{AccessToken: "other-at"})) + + s, err = m.Status(ctx) + require.NoError(t, err) + assert.Equal(t, &Status{LoggedIn: true, SessionIDs: []string{"default", "other"}, HasStoredAPIToken: true}, s) + + require.NoError(t, m.Logout(ctx, "default")) + assert.ElementsMatch(t, []string{"rt", "at", "api-at"}, srv.revoked) + ids, err := m.Store.List() + require.NoError(t, err) + assert.Equal(t, []string{"other"}, ids) +} + +func TestManager_LogoutToReplace_LockedKeychain(t *testing.T) { + keyring.MockInit() + srv := newTestAuthServer(t) + m := newTestManager(srv, t.TempDir(), &config.Auth{}) + m.Store.Service = "test-cli-auth" + m.Store.UseKeychain = true + apiSession := APITokenSessionID("good-api-token") + require.NoError(t, m.Store.Save("default", &store.Entry{APIToken: "good-api-token"})) + require.NoError(t, m.Store.Save(apiSession, &store.Entry{AccessToken: "api-at"})) + + keyring.MockInitWithError(errors.New("locked")) + var stderr strings.Builder + m.Stderr = &stderr + require.NoError(t, m.LogoutToReplace(context.Background(), "default", "good-api-token")) + assert.Contains(t, stderr.String(), "Warning: failed to load credentials in the keychain") + + // The new credentials can be saved, in files. + require.NoError(t, m.Store.Save(apiSession, &store.Entry{AccessToken: "new-api-at"})) + require.NoError(t, m.Store.Save("default", &store.Entry{APIToken: "good-api-token"})) +} + +func TestMigrator(t *testing.T) { + srv := newTestAuthServer(t) + dir := filepath.Join(t.TempDir(), "auth") + sessionDir := filepath.Join(filepath.Dir(dir), ".session") + markerPath := filepath.Join(dir, store.MigrationMarker) + exported := `{ + "default": {"access_token": "a", "refresh_token": "r", "token_type": "bearer", "expires": 123}, + "other": {"api_token": "t"}, + "api-token-abc": {"access_token": "skipped"}, + "empty": {} + }` + // The legacy files only need to exist: their content is read by the export. + for _, f := range []string{"sess-default/sess-default.json", "sess-cli-other/api-token"} { + require.NoError(t, os.MkdirAll(filepath.Dir(filepath.Join(sessionDir, f)), 0o700)) + require.NoError(t, os.WriteFile(filepath.Join(sessionDir, f), nil, 0o600)) + } + var exports, deletes int + failExport, failDelete := true, true + mg := &Migrator{Export: func(_ context.Context, del bool) ([]byte, error) { + if del { + deletes++ + if failDelete { + return nil, fmt.Errorf("delete failed") + } + return nil, nil + } + exports++ + if failExport { + return nil, fmt.Errorf("export failed") + } + return []byte(exported), nil + }} + var stderr strings.Builder + newManager := func() *Manager { + m := newTestManager(srv, dir, &config.Auth{}) + m.Migrator = mg + m.Stderr = &stderr + return m + } + // expireRetry makes a failed step due for a retry. + expireRetry := func() { + b, err := os.ReadFile(markerPath) + require.NoError(t, err) + var mk migrationMarker + require.NoError(t, json.Unmarshal(b, &mk)) + assert.Greater(t, mk.RetryAfter, time.Now().Unix()) + mk.RetryAfter = 0 + require.NoError(t, writeMarker(markerPath, &mk)) + } + + // A failed export does not block authentication, and is retried later. + ids, err := newManager().SessionIDs(context.Background()) + require.NoError(t, err) + assert.Empty(t, ids) + assert.Contains(t, stderr.String(), "failed to migrate credentials from the legacy CLI: export failed") + _, err = newManager().SessionIDs(context.Background()) + require.NoError(t, err) + assert.Equal(t, 1, exports, "the export must not be retried immediately") + + failExport = false + expireRetry() + ids, err = newManager().SessionIDs(context.Background()) + require.NoError(t, err) + assert.Equal(t, []string{"default", "other"}, ids) + e, err := newManager().Load(context.Background(), "default") + require.NoError(t, err) + assert.Equal(t, &store.Entry{AccessToken: "a", RefreshToken: "r", TokenType: "bearer", Expires: 123}, e) + assert.Equal(t, 2, exports) + assert.Equal(t, 1, deletes, "a failed delete must not be retried immediately") + + failDelete = false + expireRetry() + _, err = newManager().SessionIDs(context.Background()) + require.NoError(t, err) + _, err = newManager().SessionIDs(context.Background()) + require.NoError(t, err) + assert.Equal(t, 2, exports, "sessions must only be imported once") + assert.Equal(t, 2, deletes) + + b, err := os.ReadFile(markerPath) + require.NoError(t, err) + assert.JSONEq(t, `{}`, string(b)) +} + +func TestMigrator_SkipsUnneededExports(t *testing.T) { + cases := []struct { + name string + files []string // Legacy files, relative to the writable dir. + stored []string // Sessions in the Go store. + marker string // The marker before the run. + wantExports int + wantDeletes int + }{ + {name: "no legacy storage", wantExports: 0, wantDeletes: 0}, + {name: "first run", files: []string{".session/sess-default/sess-default.json"}, wantExports: 1, wantDeletes: 1}, + { + name: "logged in again after a failed export", + files: []string{".session/sess-default/sess-default.json"}, + stored: []string{"default"}, + marker: `{"export_pending": true, "retry_after": 9999999999}`, + wantDeletes: 1, + }, + { + name: "nothing left after a failed export", + marker: `{"export_pending": true, "retry_after": 9999999999}`, + }, + { + name: "another session still needs the export", + files: []string{".session/sess-default/sess-default.json", ".session/sess-cli-work/api-token"}, + stored: []string{"default"}, + marker: `{"export_pending": true, "retry_after": 9999999999}`, + }, + { + name: "keychain sessions may exist", + files: []string{"credential-helper"}, + marker: `{"export_pending": true, "retry_after": 9999999999}`, + }, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + writableDir := t.TempDir() + m := newTestManager(newTestAuthServer(t), filepath.Join(writableDir, "auth"), &config.Auth{}) + for _, f := range c.files { + require.NoError(t, os.MkdirAll(filepath.Dir(filepath.Join(writableDir, f)), 0o700)) + require.NoError(t, os.WriteFile(filepath.Join(writableDir, f), nil, 0o600)) + } + for _, id := range c.stored { + require.NoError(t, m.Store.Save(id, &store.Entry{AccessToken: "a"})) + } + markerPath := filepath.Join(m.Store.Dir, store.MigrationMarker) + if c.marker != "" { + require.NoError(t, store.WriteFileAtomic(markerPath, []byte(c.marker))) + } + var exports, deletes int + m.Migrator = &Migrator{Export: func(_ context.Context, del bool) ([]byte, error) { + if del { + deletes++ + } else { + exports++ + } + return []byte("{}"), nil + }} + + _, err := m.SessionIDs(context.Background()) + require.NoError(t, err) + assert.Equal(t, c.wantExports, exports) + assert.Equal(t, c.wantDeletes, deletes) + }) + } +} + +func TestTransport(t *testing.T) { + srv := newTestAuthServer(t) + srv.addRefreshToken("rt-0") + m := newTestManager(srv, t.TempDir(), &config.Auth{}) + require.NoError(t, m.Store.Save("default", &store.Entry{ + AccessToken: "revoked", RefreshToken: "rt-0", Expires: time.Now().Add(time.Hour).Unix(), + })) + + var authHeaders []string + api := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + authHeaders = append(authHeaders, r.Header.Get("Authorization")) + switch { + case r.URL.Path == "/step-up": + w.Header().Set("WWW-Authenticate", `Bearer error="insufficient_user_authentication"`) + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"amr": ["mfa"], "max_age": 60}`)) + case r.Header.Get("Authorization") != "Bearer at-1": + w.WriteHeader(http.StatusUnauthorized) + default: + _, _ = w.Write([]byte("ok")) + } + })) + defer api.Close() + + client := NewClient(m, nil) + resp, err := client.Get(api.URL) + require.NoError(t, err) + resp.Body.Close() + assert.Equal(t, http.StatusOK, resp.StatusCode) + assert.Equal(t, []string{"Bearer revoked", "Bearer at-1"}, authHeaders) + + _, err = client.Get(api.URL + "/step-up") //nolint:bodyclose // the request fails + lerr, ok := AsLoginRequired(err) + require.True(t, ok) + assert.Equal(t, []string{"mfa"}, lerr.AuthMethods) + assert.Equal(t, 60, *lerr.MaxAge) + assert.Equal(t, "Multi-factor authentication (MFA) is required.", lerr.Message()) +} diff --git a/internal/auth/migrate.go b/internal/auth/migrate.go new file mode 100644 index 000000000..016f56ea0 --- /dev/null +++ b/internal/auth/migrate.go @@ -0,0 +1,230 @@ +package auth + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io/fs" + "os" + "path/filepath" + "runtime" + "slices" + "strings" + "time" + + "github.com/upsun/cli/internal/auth/store" + "github.com/upsun/cli/internal/config" +) + +// Migrator imports sessions from the legacy CLI's storage, once. +type Migrator struct { + // Export runs the legacy CLI's hidden auth:export-sessions command and returns its output. + // With del set, it runs "auth:export-sessions --delete" instead, which deletes the exported copies. + Export func(ctx context.Context, del bool) ([]byte, error) + + // DebugLog logs a debug message. It may be nil. + DebugLog func(format string, args ...any) +} + +// migrationMarker is the content of the marker file. +type migrationMarker struct { + // ExportPending records that the export failed, and DeletePending that the legacy CLI's copies were not deleted. + ExportPending bool `json:"export_pending,omitempty"` + DeletePending bool `json:"delete_pending,omitempty"` + // RetryAfter is when a failed step may be retried, as a Unix timestamp. + RetryAfter int64 `json:"retry_after,omitempty"` +} + +// migrationRetryDelay is how long to wait before retrying a failed migration step. +const migrationRetryDelay = time.Hour + +// exportedSession is a session as exported by the legacy CLI. +type exportedSession struct { + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token"` + TokenType string `json:"token_type"` + Expires int64 `json:"expires"` + APIToken string `json:"api_token"` +} + +// Run migrates sessions if that has not been done yet. +// +// Sessions are never imported twice, because a second import would overwrite fresh tokens with rotated ones. +// A failure does not block authentication: it is reported, and retried after migrationRetryDelay. +func (mg *Migrator) Run(ctx context.Context, m *Manager) error { + markerPath := filepath.Join(m.Store.Dir, store.MigrationMarker) + if mk, err := readMarker(markerPath); err != nil || !mg.due(m, mk) { + return err + } + if !m.Settings.DisableLocks { + unlock, err := m.fileLock(ctx, filepath.Join(m.Store.Dir, ".migrate.lock")) + if err != nil { + return err + } + defer unlock() + } + mk, err := readMarker(markerPath) + if err != nil || !mg.due(m, mk) { + return err + } + + if mk == nil || mk.ExportPending { + switch legacy := mg.legacyState(m); { + case legacy == legacyEmpty: + return writeMarker(markerPath, &migrationMarker{}) + case mk != nil && legacy == legacyImported: + // Only the deletion is left. + default: + if err := mg.importSessions(ctx, m); err != nil { + if m.Stderr != nil { + fmt.Fprintf(m.Stderr, "Warning: failed to migrate credentials from the legacy CLI: %s\n", err) + } + return writeMarker(markerPath, mg.retryLater(&migrationMarker{ExportPending: true})) + } + } + } + + // The legacy CLI's copies are deleted to avoid having two sources of truth. + if _, err := mg.Export(ctx, true); err != nil { + mg.debugf("Failed to delete the legacy CLI's credentials: %s", err) + return writeMarker(markerPath, mg.retryLater(&migrationMarker{DeletePending: true})) + } + return writeMarker(markerPath, &migrationMarker{}) +} + +// due reports whether a migration step should run, given the marker (nil if there is none). +// +// A failed export is retried early if it is no longer needed, e.g. because the user logged in again. +func (mg *Migrator) due(m *Manager, mk *migrationMarker) bool { + switch { + case mk == nil: + return true + case mk.ExportPending && mg.legacyState(m) != legacyUnknown: + return true + default: + return (mk.ExportPending || mk.DeletePending) && time.Now().Unix() >= mk.RetryAfter + } +} + +type legacyStorageState int + +const ( + legacyUnknown legacyStorageState = iota // Sessions may exist that are not in the Go store. + legacyImported // Every legacy session is already in the Go store. + legacyEmpty // The legacy storage holds no sessions. +) + +// legacyState checks the legacy CLI's session files by their names, without reading them. +// +// Sessions in the keychain can only be listed by the legacy credential helper, so they count as unknown. +func (mg *Migrator) legacyState(m *Manager) legacyStorageState { + writableDir := filepath.Dir(m.Store.Dir) + helper := filepath.Join(writableDir, "credential-helper") + if runtime.GOOS == "windows" { + helper += ".exe" + } + if _, err := os.Stat(helper); err == nil { + return legacyUnknown + } + sessionDir := filepath.Join(writableDir, ".session") + var ids []string + files, _ := filepath.Glob(filepath.Join(sessionDir, "sess-*", "sess-*.json")) + for _, f := range files { + id := strings.TrimPrefix(strings.TrimSuffix(filepath.Base(f), ".json"), "sess-") + if filepath.Base(filepath.Dir(f)) == "sess-"+id { + ids = append(ids, id) + } + } + tokenFiles, _ := filepath.Glob(filepath.Join(sessionDir, "sess-cli-*", "api-token")) + for _, f := range tokenFiles { + ids = append(ids, strings.TrimPrefix(filepath.Base(filepath.Dir(f)), "sess-cli-")) + } + if len(ids) == 0 { + return legacyEmpty + } + stored, err := m.Store.List() + if err != nil { + return legacyUnknown + } + for _, id := range ids { + if !strings.HasPrefix(id, "api-token-") && config.ValidateSessionID(id) == nil && !slices.Contains(stored, id) { + return legacyUnknown + } + } + return legacyImported +} + +func (mg *Migrator) retryLater(mk *migrationMarker) *migrationMarker { + mk.RetryAfter = time.Now().Add(migrationRetryDelay).Unix() + return mk +} + +// importSessions exports sessions from the legacy CLI and saves the ones that are not already stored. +func (mg *Migrator) importSessions(ctx context.Context, m *Manager) error { + out, err := mg.Export(ctx, false) + if err != nil { + return err + } + var sessions map[string]exportedSession + if err := json.Unmarshal(out, &sessions); err != nil { + return fmt.Errorf("invalid export: %w", err) + } + for id, s := range sessions { + // API token sessions are skipped, as the token is exchanged again. + if strings.HasPrefix(id, "api-token-") || config.ValidateSessionID(id) != nil { + continue + } + if s.AccessToken == "" && s.RefreshToken == "" && s.APIToken == "" { + continue + } + if existing, err := m.Store.Load(id); err != nil { + return err + } else if existing != nil { + continue + } + entry := &store.Entry{ + AccessToken: s.AccessToken, + RefreshToken: s.RefreshToken, + TokenType: s.TokenType, + Expires: s.Expires, + APIToken: s.APIToken, + } + if err := m.Store.Save(id, entry); err != nil { + return err + } + mg.debugf("Migrated session: %s", id) + } + return nil +} + +func (mg *Migrator) debugf(format string, args ...any) { + if mg.DebugLog != nil { + mg.DebugLog(format, args...) + } +} + +func readMarker(path string) (*migrationMarker, error) { + b, err := os.ReadFile(path) + if err != nil { + if errors.Is(err, fs.ErrNotExist) { + return nil, nil + } + return nil, err + } + mk := &migrationMarker{} + if len(b) > 0 { + if err := json.Unmarshal(b, mk); err != nil { + return nil, fmt.Errorf("invalid file %s: %w", path, err) + } + } + return mk, nil +} + +func writeMarker(path string, mk *migrationMarker) error { + b, err := json.Marshal(mk) + if err != nil { + return err + } + return store.WriteFileAtomic(path, b) +} diff --git a/internal/auth/oauth.go b/internal/auth/oauth.go new file mode 100644 index 000000000..84f79c607 --- /dev/null +++ b/internal/auth/oauth.go @@ -0,0 +1,228 @@ +package auth + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptrace" + "net/url" + "strings" + "time" + + "github.com/upsun/cli/internal/auth/store" +) + +// OAuthError is an OAuth 2.0 error response (RFC 6749, section 5.2). +type OAuthError struct { + StatusCode int `json:"-"` + Code string `json:"error"` + Description string `json:"error_description"` + Hint string `json:"error_hint"` +} + +func (e *OAuthError) Error() string { + if e.Description != "" { + return e.Description + } + if e.Code != "" { + return e.Code + } + return fmt.Sprintf("OAuth 2.0 request failed with status %d", e.StatusCode) +} + +// tokenResponse is a successful token endpoint response. +type tokenResponse struct { + AccessToken string `json:"access_token"` + TokenType string `json:"token_type"` + RefreshToken string `json:"refresh_token"` + ExpiresIn int64 `json:"expires_in"` +} + +// OAuthClient calls the OAuth 2.0 token and revocation endpoints, as a public client. +type OAuthClient struct { + HTTPClient *http.Client + TokenURL string + RevokeURL string + ClientID string + + // retryDelay is the base delay between retries, for tests. + retryDelay time.Duration +} + +// requestTimeout bounds each request to the auth server. +const requestTimeout = 30 * time.Second + +// ExchangeCode exchanges an authorization code (with its PKCE verifier) for tokens. +func (c *OAuthClient) ExchangeCode(ctx context.Context, code, verifier, redirectURI string) (*store.Entry, error) { + form := url.Values{ + "grant_type": {"authorization_code"}, + "code": {code}, + "redirect_uri": {redirectURI}, + "code_verifier": {verifier}, + } + return c.postToken(ctx, form, true) +} + +// ExchangeAPIToken exchanges an API token for tokens, using the "api_token" grant. +func (c *OAuthClient) ExchangeAPIToken(ctx context.Context, apiToken string) (*store.Entry, error) { + return c.withRetries(ctx, 1, func(ctx context.Context) (*store.Entry, bool, error) { + return c.postTokenTraced(ctx, c.clientForm(url.Values{"grant_type": {"api_token"}, "api_token": {apiToken}})) + }) +} + +// Refresh uses a refresh token to get new tokens. +// +// Connection errors before the request is sent are retried twice. Errors after it is sent are not retried, because +// the server may have consumed the refresh token. +func (c *OAuthClient) Refresh(ctx context.Context, refreshToken string) (*store.Entry, error) { + return c.withRetries(ctx, 0, func(ctx context.Context) (*store.Entry, bool, error) { + form := c.clientForm(url.Values{"grant_type": {"refresh_token"}, "refresh_token": {refreshToken}}) + return c.postTokenTraced(ctx, form) + }) +} + +// Revoke revokes a token. The hint is "access_token" or "refresh_token". +func (c *OAuthClient) Revoke(ctx context.Context, token, hint string) error { + form := c.clientForm(url.Values{"token": {token}, "token_type_hint": {hint}}) + var err error + for attempt := range 2 { + var status int + status, _, err = c.post(ctx, c.RevokeURL, form, false) + if err != nil { + return err + } + if status < 300 { + return nil + } + err = fmt.Errorf("token revocation failed with status %d", status) + // Retry once on a retry status, as the legacy CLI does. + switch status { + case 408, 429, 502, 503, 504: + if attempt == 0 { + continue + } + } + return err + } + return err +} + +func (c *OAuthClient) clientForm(form url.Values) url.Values { + form.Set("client_id", c.ClientID) + form.Set("client_secret", "") + return form +} + +// withRetries runs a token request, retrying transient failures. The callback reports whether the request was sent. +// Transient failures after the request was sent are retried up to sentRetries times. +func (c *OAuthClient) withRetries( + ctx context.Context, + sentRetries int, + fn func(ctx context.Context) (*store.Entry, bool, error), +) (*store.Entry, error) { + delay := c.retryDelay + if delay == 0 { + delay = 500 * time.Millisecond + } + unsentRetries := 2 + for { + e, sent, err := fn(ctx) + if err == nil || ctx.Err() != nil { + return e, err + } + var oerr *OAuthError + isOAuthError := errors.As(err, &oerr) + switch { + case !sent && unsentRetries > 0: + unsentRetries-- + case sent && sentRetries > 0 && (!isOAuthError || oerr.StatusCode >= 500): + sentRetries-- + default: + return nil, err + } + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-time.After(delay): + } + delay *= 2 + } +} + +// postTokenTraced posts to the token endpoint, and reports whether the request was written. +func (c *OAuthClient) postTokenTraced(ctx context.Context, form url.Values) (*store.Entry, bool, error) { + var sent bool + trace := &httptrace.ClientTrace{WroteRequest: func(httptrace.WroteRequestInfo) { sent = true }} + e, err := c.postToken(httptrace.WithClientTrace(ctx, trace), form, false) + var oerr *OAuthError + if errors.As(err, &oerr) { + // A response was received. + sent = true + } + return e, sent, err +} + +func (c *OAuthClient) postToken(ctx context.Context, form url.Values, basicAuth bool) (*store.Entry, error) { + status, body, err := c.post(ctx, c.TokenURL, form, basicAuth) + if err != nil { + return nil, err + } + if status >= 300 { + oerr := &OAuthError{StatusCode: status} + _ = json.Unmarshal(body, oerr) + return nil, oerr + } + var tr tokenResponse + if err := json.Unmarshal(body, &tr); err != nil { + return nil, fmt.Errorf("invalid token response: %w", err) + } + if tr.AccessToken == "" { + oerr := &OAuthError{StatusCode: status} + if json.Unmarshal(body, oerr) == nil && oerr.Code != "" { + return nil, oerr + } + return nil, errors.New("invalid token response: no access token") + } + e := &store.Entry{ + AccessToken: tr.AccessToken, + RefreshToken: tr.RefreshToken, + TokenType: tr.TokenType, + } + if tr.ExpiresIn > 0 { + e.Expires = time.Now().Unix() + tr.ExpiresIn + } else if exp, err := unsafeGetJWTExpiry(tr.AccessToken); err == nil { + e.Expires = exp.Unix() + } + return e, nil +} + +// post sends a form, and returns the response status and body. Each request is bounded by requestTimeout. +func (c *OAuthClient) post( + ctx context.Context, u string, form url.Values, basicAuth bool, +) (status int, body []byte, err error) { + ctx, cancel := context.WithTimeout(ctx, requestTimeout) + defer cancel() + req, err := http.NewRequestWithContext(ctx, http.MethodPost, u, strings.NewReader(form.Encode())) + if err != nil { + return 0, nil, err + } + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + req.Header.Set("Accept", "application/json") + if basicAuth { + req.SetBasicAuth(c.ClientID, "") + } + hc := c.HTTPClient + if hc == nil { + hc = http.DefaultClient + } + resp, err := hc.Do(req) + if err != nil { + return 0, nil, err + } + defer resp.Body.Close() + body, err = io.ReadAll(io.LimitReader(resp.Body, 1<<20)) + return resp.StatusCode, body, err +} diff --git a/internal/auth/store/keychain.go b/internal/auth/store/keychain.go new file mode 100644 index 000000000..928679ce6 --- /dev/null +++ b/internal/auth/store/keychain.go @@ -0,0 +1,36 @@ +package store + +import ( + "os" + "runtime" + "strings" +) + +// KeychainSupported reports whether the system keychain is likely to work. +// +// On Linux this keeps the legacy CLI's conditions: a display, a GNOME session, not in a snap or a container, and a +// Secret Service on D-Bus. +func KeychainSupported() bool { + switch runtime.GOOS { + case "darwin", "windows": + return true + case "linux": + if d := os.Getenv("DISPLAY"); d == "" || d == "none" { + return false + } + if !strings.Contains(strings.ToUpper(os.Getenv("XDG_CURRENT_DESKTOP")), "GNOME") { + return false + } + for _, v := range []string{"SNAP_CONTEXT", "container", "DOCKER_IP"} { + if _, ok := os.LookupEnv(v); ok { + return false + } + } + if _, err := os.Stat("/.dockerenv"); err == nil { + return false + } + return secretServiceAvailable() + default: + return false + } +} diff --git a/internal/auth/store/keychain_linux.go b/internal/auth/store/keychain_linux.go new file mode 100644 index 000000000..fc55c1753 --- /dev/null +++ b/internal/auth/store/keychain_linux.go @@ -0,0 +1,34 @@ +package store + +import ( + "slices" + + "github.com/godbus/dbus/v5" +) + +const secretServiceName = "org.freedesktop.secrets" + +// secretServiceAvailable checks whether the Secret Service D-Bus name is owned or can be activated. +func secretServiceAvailable() bool { + conn, err := dbus.SessionBusPrivate() + if err != nil { + return false + } + defer conn.Close() + if err := conn.Auth(nil); err != nil { + return false + } + if err := conn.Hello(); err != nil { + return false + } + var hasOwner bool + call := conn.BusObject().Call("org.freedesktop.DBus.NameHasOwner", 0, secretServiceName) + if err := call.Store(&hasOwner); err == nil && hasOwner { + return true + } + var activatable []string + if err := conn.BusObject().Call("org.freedesktop.DBus.ListActivatableNames", 0).Store(&activatable); err != nil { + return false + } + return slices.Contains(activatable, secretServiceName) +} diff --git a/internal/auth/store/keychain_other.go b/internal/auth/store/keychain_other.go new file mode 100644 index 000000000..c0389e042 --- /dev/null +++ b/internal/auth/store/keychain_other.go @@ -0,0 +1,7 @@ +//go:build !linux + +package store + +func secretServiceAvailable() bool { + return false +} diff --git a/internal/auth/store/store.go b/internal/auth/store/store.go new file mode 100644 index 000000000..993a44e32 --- /dev/null +++ b/internal/auth/store/store.go @@ -0,0 +1,313 @@ +// Package store saves OAuth 2.0 credentials, one entry per session ID, in the system keychain or in files. +package store + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "io/fs" + "os" + "path/filepath" + "runtime" + "slices" + "strings" + "time" + + "github.com/zalando/go-keyring" +) + +// Entry holds the credentials for one session. +type Entry struct { + AccessToken string `json:"access_token,omitempty"` + RefreshToken string `json:"refresh_token,omitempty"` + TokenType string `json:"token_type,omitempty"` + Expires int64 `json:"expires,omitempty"` // A Unix timestamp. + + // APIToken is set for a token saved by auth:api-token-login. + APIToken string `json:"api_token,omitempty"` +} + +// Backend is where an entry's secrets are stored. +type Backend string + +const ( + BackendKeychain Backend = "keychain" + BackendFile Backend = "file" +) + +// sessionFile is the content of /.json. It holds no secrets in keychain mode. +type sessionFile struct { + Backend Backend `json:"backend"` + Entry *Entry `json:"entry,omitempty"` +} + +// MigrationMarker is the name of the file that records the migration from the legacy CLI's storage. +const MigrationMarker = ".migrated" + +const defaultKeychainTimeout = 10 * time.Second + +// Store saves entries. The backend is chosen when a session is first saved, and recorded in the session file. +type Store struct { + // Dir is the directory for session files and locks, e.g. ~/.upsun-cli/auth. + Dir string + // Service is the keychain service name. + Service string + // UseKeychain reports whether a new session should try the keychain. + UseKeychain bool + // KeychainTimeout limits each keychain read, and the first write of a session. It defaults to 10s. + // Other changes are waited for, so that they cannot finish after the session's lock is released. + KeychainTimeout time.Duration + // Stderr receives a notice while waiting for the keychain. It may be nil. + Stderr io.Writer +} + +// KeychainError is returned when the keychain cannot be used for a session that is stored there. +type KeychainError struct { + Op string + Err error +} + +func (e *KeychainError) Error() string { + return fmt.Sprintf("failed to %s credentials in the keychain: %s\n"+ + "Check that the keychain is unlocked, or log in again to store credentials in a file instead.", e.Op, e.Err) +} + +func (e *KeychainError) Unwrap() error { return e.Err } + +// Load returns the entry for a session, or nil if there is none. +func (s *Store) Load(id string) (*Entry, error) { + sf, err := s.readSessionFile(id) + if err != nil || sf == nil { + return nil, err + } + if sf.Backend != BackendKeychain { + return sf.Entry, nil + } + secret, err := s.keychain(false, func() (string, error) { return keyring.Get(s.Service, id) }) + if errors.Is(err, keyring.ErrNotFound) { + return nil, nil + } + if err != nil { + return nil, &KeychainError{Op: "load", Err: err} + } + var e Entry + if err := json.Unmarshal([]byte(secret), &e); err != nil { + return nil, fmt.Errorf("invalid credentials in the keychain: %w", err) + } + return &e, nil +} + +// Save saves the entry for a session. +func (s *Store) Save(id string, e *Entry) error { + sf, err := s.readSessionFile(id) + if err != nil { + return err + } + b, err := json.Marshal(e) //nolint:gosec // the entry is stored in the keychain + if err != nil { + return err + } + backend := BackendFile + switch { + case sf != nil: + backend = sf.Backend + case s.UseKeychain: + // The backend is chosen once, so any keychain failure here (including data that is too big) falls back to + // a file. + if err := s.keychainSet(id, b, false); err == nil { + backend = BackendKeychain + } + } + if backend == BackendKeychain { + if sf != nil { + if err := s.keychainSet(id, b, true); err != nil { + return &KeychainError{Op: "save", Err: err} + } + } + return s.writeSessionFile(id, &sessionFile{Backend: BackendKeychain}) + } + return s.writeSessionFile(id, &sessionFile{Backend: BackendFile, Entry: e}) +} + +// Delete removes the entry for a session. It does nothing if there is none. +func (s *Store) Delete(id string) error { + sf, err := s.readSessionFile(id) + if err != nil { + return err + } + if sf == nil { + return nil + } + if sf.Backend == BackendKeychain { + _, err := s.keychain(true, func() (string, error) { return "", keyring.Delete(s.Service, id) }) + if err != nil && !errors.Is(err, keyring.ErrNotFound) { + return &KeychainError{Op: "delete", Err: err} + } + } + if err := os.Remove(s.sessionFilePath(id)); err != nil && !errors.Is(err, fs.ErrNotExist) { + return err + } + return nil +} + +// List returns the IDs of all stored sessions, sorted. +func (s *Store) List() ([]string, error) { + entries, err := os.ReadDir(s.Dir) + if err != nil { + if errors.Is(err, fs.ErrNotExist) { + return nil, nil + } + return nil, err + } + var ids []string + for _, e := range entries { + if id, ok := strings.CutSuffix(e.Name(), ".json"); ok && !e.IsDir() && !strings.HasPrefix(id, ".") { + ids = append(ids, id) + } + } + slices.Sort(ids) + return ids, nil +} + +// DeleteAll removes every session, and all other files except locks and the migration marker. +// +// Lock files are kept, as another process may hold a lock on them, and so are the files of sessions whose secrets +// could not be deleted. +func (s *Store) DeleteAll() error { + ids, err := s.List() + if err != nil { + return err + } + var errs []error + // The files of sessions that could not be deleted are kept, so that deleting their secrets can be retried. + keep := map[string]bool{MigrationMarker: true} + for _, id := range ids { + if err := s.Delete(id); err != nil { + errs = append(errs, err) + keep[id+".json"] = true + } + } + entries, err := os.ReadDir(s.Dir) + if err != nil && !errors.Is(err, fs.ErrNotExist) { + errs = append(errs, err) + } + for _, e := range entries { + if !keep[e.Name()] && !strings.HasSuffix(e.Name(), ".lock") { + errs = append(errs, os.RemoveAll(filepath.Join(s.Dir, e.Name()))) + } + } + return errors.Join(errs...) +} + +// Forget removes a session's file without deleting its secrets, e.g. if they are in a keychain that cannot be used. +func (s *Store) Forget(id string) error { + if err := os.Remove(s.sessionFilePath(id)); err != nil && !errors.Is(err, fs.ErrNotExist) { + return err + } + return nil +} + +// LockPath returns the path of the lock file for a session. +func (s *Store) LockPath(id string) string { + return filepath.Join(s.Dir, id+".lock") +} + +func (s *Store) sessionFilePath(id string) string { + return filepath.Join(s.Dir, id+".json") +} + +func (s *Store) readSessionFile(id string) (*sessionFile, error) { + b, err := os.ReadFile(s.sessionFilePath(id)) + if err != nil { + if errors.Is(err, fs.ErrNotExist) { + return nil, nil + } + return nil, err + } + var sf sessionFile + if err := json.Unmarshal(b, &sf); err != nil { + return nil, fmt.Errorf("invalid session file %s: %w", s.sessionFilePath(id), err) + } + return &sf, nil +} + +func (s *Store) writeSessionFile(id string, sf *sessionFile) error { + b, err := json.Marshal(sf) + if err != nil { + return err + } + return WriteFileAtomic(s.sessionFilePath(id), b) +} + +func (s *Store) keychainSet(id string, secret []byte, wait bool) error { + _, err := s.keychain(wait, func() (string, error) { return "", keyring.Set(s.Service, id, string(secret)) }) + return err +} + +// keychain runs a keychain call with a timeout. If wait is set, a notice is printed at the timeout, and the call is +// still waited for. +func (s *Store) keychain(wait bool, fn func() (string, error)) (string, error) { + timeout := s.KeychainTimeout + if timeout == 0 { + timeout = defaultKeychainTimeout + } + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + type result struct { + v string + err error + } + ch := make(chan result, 1) + go func() { + v, err := fn() + ch <- result{v, err} + }() + select { + case r := <-ch: + return r.v, r.err + case <-ctx.Done(): + if !wait { + return "", fmt.Errorf("timed out after %s", timeout) + } + } + if s.Stderr != nil { + fmt.Fprintln(s.Stderr, "Waiting for the keychain. Check whether it needs to be unlocked.") + } + r := <-ch + return r.v, r.err +} + +// WriteFileAtomic writes a file with 0600 permissions via a synced temporary file, creating the directory (0700). +func WriteFileAtomic(path string, data []byte) error { + dir := filepath.Dir(path) + if err := os.MkdirAll(dir, 0o700); err != nil { + return err + } + f, err := os.CreateTemp(dir, ".tmp-"+filepath.Base(path)+"-*") + if err != nil { + return err + } + tmp := f.Name() + defer os.Remove(tmp) + if _, err := f.Write(data); err != nil { + _ = f.Close() + return err + } + if err := f.Sync(); err != nil { + _ = f.Close() + return err + } + if err := f.Close(); err != nil { + return err + } + // On Windows, replacing a file fails while another process has it open, e.g. a reader outside the lock. + for attempt := 0; ; attempt++ { + err = os.Rename(tmp, path) + if err == nil || runtime.GOOS != "windows" || attempt == 40 { + return err + } + time.Sleep(25 * time.Millisecond) + } +} diff --git a/internal/auth/store/store_test.go b/internal/auth/store/store_test.go new file mode 100644 index 000000000..2074687f1 --- /dev/null +++ b/internal/auth/store/store_test.go @@ -0,0 +1,168 @@ +package store + +import ( + "errors" + "os" + "path/filepath" + "runtime" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/zalando/go-keyring" +) + +func TestStore(t *testing.T) { + cases := []struct { + name string + useKeychain bool + keyringErr error + wantBackend Backend + }{ + {name: "file", wantBackend: BackendFile}, + {name: "keychain", useKeychain: true, wantBackend: BackendKeychain}, + {name: "keychain unavailable", useKeychain: true, keyringErr: errors.New("locked"), wantBackend: BackendFile}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + if c.keyringErr != nil { + keyring.MockInitWithError(c.keyringErr) + } else { + keyring.MockInit() + } + s := &Store{Dir: filepath.Join(t.TempDir(), "auth"), Service: "test-cli-auth", UseKeychain: c.useKeychain} + + e, err := s.Load("default") + require.NoError(t, err) + assert.Nil(t, e) + + entry := &Entry{AccessToken: "a", RefreshToken: "r", TokenType: "bearer", Expires: 123} + require.NoError(t, s.Save("default", entry)) + + sf, err := s.readSessionFile("default") + require.NoError(t, err) + assert.Equal(t, c.wantBackend, sf.Backend) + if c.wantBackend == BackendKeychain { + assert.Nil(t, sf.Entry, "the session file must not hold secrets in keychain mode") + } + if runtime.GOOS != "windows" { + info, err := os.Stat(s.sessionFilePath("default")) + require.NoError(t, err) + assert.Equal(t, os.FileMode(0o600), info.Mode().Perm()) + } + + loaded, err := s.Load("default") + require.NoError(t, err) + assert.Equal(t, entry, loaded) + + entry.AccessToken = "a2" + require.NoError(t, s.Save("default", entry)) + loaded, err = s.Load("default") + require.NoError(t, err) + assert.Equal(t, "a2", loaded.AccessToken) + + require.NoError(t, s.Save("other", &Entry{APIToken: "t"})) + ids, err := s.List() + require.NoError(t, err) + assert.Equal(t, []string{"default", "other"}, ids) + + require.NoError(t, s.Delete("default")) + loaded, err = s.Load("default") + require.NoError(t, err) + assert.Nil(t, loaded) + require.NoError(t, s.Delete("default")) + }) + } +} + +func TestStore_KeychainFailsLater(t *testing.T) { + keyring.MockInit() + s := &Store{Dir: t.TempDir(), Service: "test-cli-auth", UseKeychain: true} + require.NoError(t, s.Save("default", &Entry{AccessToken: "a"})) + + // A session stored in the keychain must not silently move to a file. + keyring.MockInitWithError(errors.New("locked")) + err := s.Save("default", &Entry{AccessToken: "b"}) + var kerr *KeychainError + require.ErrorAs(t, err, &kerr) + _, err = s.Load("default") + require.ErrorAs(t, err, &kerr) +} + +func TestStore_KeychainTooBig(t *testing.T) { + keyring.MockInitWithError(keyring.ErrSetDataTooBig) + s := &Store{Dir: t.TempDir(), Service: "test-cli-auth", UseKeychain: true} + require.NoError(t, s.Save("default", &Entry{AccessToken: "a"})) + sf, err := s.readSessionFile("default") + require.NoError(t, err) + assert.Equal(t, BackendFile, sf.Backend) +} + +func TestStore_KeychainTimeout(t *testing.T) { + s := &Store{KeychainTimeout: 10 * time.Millisecond} + _, err := s.keychain(false, func() (string, error) { + time.Sleep(time.Second) + return "", nil + }) + assert.ErrorContains(t, err, "timed out") +} + +func TestStore_KeychainWait(t *testing.T) { + var stderr strings.Builder + s := &Store{KeychainTimeout: 10 * time.Millisecond, Stderr: &stderr} + done := false + v, err := s.keychain(true, func() (string, error) { + time.Sleep(100 * time.Millisecond) + done = true + return "v", nil + }) + require.NoError(t, err) + assert.Equal(t, "v", v) + assert.True(t, done, "a change must finish before the call returns") + assert.Contains(t, stderr.String(), "Waiting for the keychain") +} + +func TestStore_DeleteAll_KeychainError(t *testing.T) { + keyring.MockInit() + s := &Store{Dir: t.TempDir(), Service: "test-cli-auth", UseKeychain: true} + require.NoError(t, s.Save("a", &Entry{AccessToken: "a"})) + + keyring.MockInitWithError(errors.New("locked")) + assert.Error(t, s.DeleteAll()) + // The session file is kept, so that deleting the secret can be retried. + assert.FileExists(t, s.sessionFilePath("a")) +} + +func TestStore_DeleteAll(t *testing.T) { + keyring.MockInit() + s := &Store{Dir: t.TempDir(), Service: "test-cli-auth"} + require.NoError(t, s.Save("a", &Entry{AccessToken: "a"})) + require.NoError(t, s.Save("b", &Entry{AccessToken: "b"})) + require.NoError(t, os.WriteFile(s.LockPath("a"), nil, 0o600)) + require.NoError(t, os.WriteFile(filepath.Join(s.Dir, MigrationMarker), nil, 0o600)) + + require.NoError(t, s.DeleteAll()) + entries, err := os.ReadDir(s.Dir) + require.NoError(t, err) + names := make([]string, 0, len(entries)) + for _, e := range entries { + names = append(names, e.Name()) + } + assert.Equal(t, []string{MigrationMarker, "a.lock"}, names) +} + +func TestStore_Forget(t *testing.T) { + keyring.MockInit() + s := &Store{Dir: t.TempDir(), Service: "test-cli-auth", UseKeychain: true} + require.NoError(t, s.Save("default", &Entry{AccessToken: "a"})) + + // A session whose keychain is unusable can be replaced, e.g. with a file. + keyring.MockInitWithError(errors.New("locked")) + require.NoError(t, s.Forget("default")) + require.NoError(t, s.Save("default", &Entry{AccessToken: "b"})) + sf, err := s.readSessionFile("default") + require.NoError(t, err) + assert.Equal(t, BackendFile, sf.Backend) +} diff --git a/internal/auth/transport.go b/internal/auth/transport.go index 3c2a169e2..2e941f492 100644 --- a/internal/auth/transport.go +++ b/internal/auth/transport.go @@ -3,93 +3,101 @@ package auth import ( "bytes" "context" - "fmt" + "errors" "io" "net/http" ) -type refresher interface { - refreshToken() error - invalidateToken() error -} - -// Transport is an HTTP RoundTripper similar to golang.org/x/oauth2.Transport. -// It injects Authorization headers using a savingSource and, on a 401 response, -// clears the cached token and retries the request once. +// Transport is an HTTP RoundTripper that adds an access token to requests. +// +// On a 401 response it refreshes the token and retries the request once. A step-up authentication challenge +// (RFC 9470) returns a *LoginRequiredError instead. type Transport struct { - // base is the underlying oauth2.Transport that adds the Authorization header. - base http.RoundTripper - - // refresher is the savingSource used as the TokenSource for base; kept private - // so we can clear its cached token on 401. - refresher refresher - - LogFunc func(msg string, args ...any) + Base http.RoundTripper + Manager *Manager } -// RoundTrip adds Authorization via the underlying oauth2.Transport. If the -// response is 401 Unauthorized, it clears the cached token and retries once. func (t *Transport) RoundTrip(req *http.Request) (*http.Response, error) { - req.Body = wrapReader(req.Body) - - resp, err := t.base.RoundTrip(req) + ctx := req.Context() + body, err := bufferBody(req) + if err != nil { + return nil, err + } + tok, err := t.Manager.Token(ctx, "") + if err != nil { + return nil, err + } + resp, err := t.base().RoundTrip(withToken(req, tok, body)) + if err != nil || resp.StatusCode != http.StatusUnauthorized { + return resp, err + } + if IsStepUpChallenge(resp) { + defer resp.Body.Close() + hasAPIToken, _ := t.Manager.HasAPIToken(ctx) + return nil, StepUpError(resp, hasAPIToken) + } + flush(resp.Body) + tok, err = t.Manager.Token(ctx, tok.AccessToken) + if err != nil { + return nil, err + } + return t.base().RoundTrip(withToken(req, tok, body)) +} - // Retry on 401 - if resp != nil && resp.StatusCode == http.StatusUnauthorized { - _ = t.log("The access token needs to be refreshed. Retrying request.") - if err := t.refresher.invalidateToken(); err != nil { - return nil, fmt.Errorf("failed to invalidate token: %w", err) - } - flushReader(resp.Body) - resp, err = t.base.RoundTrip(req) +func (t *Transport) base() http.RoundTripper { + if t.Base != nil { + return t.Base } + return http.DefaultTransport +} - return resp, err +func withToken(req *http.Request, tok *Token, body []byte) *http.Request { + r := req.Clone(req.Context()) + r.Header.Set("Authorization", "Bearer "+tok.AccessToken) + if body != nil { + r.Body = io.NopCloser(bytes.NewReader(body)) + } + return r } -func (t *Transport) log(msg string, args ...any) error { - if t.LogFunc == nil { - return nil +// bufferBody reads the request body so that it can be sent twice. +func bufferBody(req *http.Request) ([]byte, error) { + if req.Body == nil { + return nil, nil } - t.LogFunc(msg, args...) - return nil + b, err := io.ReadAll(req.Body) + _ = req.Body.Close() + return b, err } -// context key for storing a custom RoundTripper. -type transportCtxKey struct{} +func flush(r io.ReadCloser) { + _, _ = io.Copy(io.Discard, r) + _ = r.Close() +} -// WithTransport returns a new context that carries the provided RoundTripper. -func WithTransport(ctx context.Context, rt http.RoundTripper) context.Context { - return context.WithValue(ctx, transportCtxKey{}, rt) +// NewClient returns an HTTP client that authenticates requests. +func NewClient(m *Manager, base http.RoundTripper) *http.Client { + return &http.Client{Transport: &Transport{Base: base, Manager: m}} } -// TransportFromContext retrieves a RoundTripper previously stored with -// WithTransport. It returns (nil, false) if none is set. -func TransportFromContext(ctx context.Context) (http.RoundTripper, bool) { - v := ctx.Value(transportCtxKey{}) - if v == nil { - return nil, false - } - rt, ok := v.(http.RoundTripper) - if !ok || rt == nil { - return nil, false - } - return rt, true +// EnsureAuthenticated checks that a token is available, refreshing it if needed. +func (m *Manager) EnsureAuthenticated(ctx context.Context) error { + _, err := m.Token(ctx, "") + return err } -func wrapReader(r io.ReadCloser) io.ReadCloser { - if r == nil { - return nil +// HasAPIToken reports whether an API token is used, whether stored or set via config. +func (m *Manager) HasAPIToken(ctx context.Context) (bool, error) { + t, err := m.apiToken(ctx) + if err != nil { + return false, err } - bodyBytes, _ := io.ReadAll(r) - _ = r.Close() - return io.NopCloser(bytes.NewBuffer(bodyBytes)) + return t != "" || m.Settings.AccessToken != "", nil } -func flushReader(r io.ReadCloser) { - if r == nil { - return - } - _, _ = io.Copy(io.Discard, r) - _ = r.Close() +// AsLoginRequired returns a *LoginRequiredError found in err. +func AsLoginRequired(err error) (*LoginRequiredError, bool) { + var lerr *LoginRequiredError + ok := errors.As(err, &lerr) + return lerr, ok } diff --git a/internal/auth/transport_test.go b/internal/auth/transport_test.go deleted file mode 100644 index 859c03cf2..000000000 --- a/internal/auth/transport_test.go +++ /dev/null @@ -1,118 +0,0 @@ -package auth - -import ( - "bytes" - "io" - "net/http" - "net/http/httptest" - "testing" - "time" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - "golang.org/x/oauth2" -) - -// mockRefresher implements both the refresher and oauth2.TokenSource interfaces for testing -type mockRefresher struct { - token *oauth2.Token -} - -func (m *mockRefresher) refreshToken() error { - m.token = &oauth2.Token{ - AccessToken: "valid", - TokenType: "Bearer", - Expiry: time.Now().Add(time.Hour), - } - return nil -} - -func (m *mockRefresher) invalidateToken() error { - m.token = &oauth2.Token{ - AccessToken: "", - TokenType: "Bearer", - Expiry: time.Now().Add(-time.Hour), - } - - return nil -} - -func (m *mockRefresher) Token() (*oauth2.Token, error) { - if m.token == nil || !m.token.Valid() { - if err := m.refreshToken(); err != nil { - return nil, err - } - } - return m.token, nil -} - -func TestTransport_RoundTrip_RetryOn401(t *testing.T) { - // Create a mock server that initially returns 401, then 200 - responseCodes := []int{} - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Read and validate the request body - body, err := io.ReadAll(r.Body) - require.NoError(t, err) - - // Check that we have the expected POST body - assert.Equal(t, "test-body-content", string(body)) - - if r.Header.Get("Authorization") != "Bearer valid" { - w.WriteHeader(http.StatusUnauthorized) - if _, err := w.Write([]byte(`{"error": "unauthorized"}`)); err != nil { - require.NoError(t, err) - } - responseCodes = append(responseCodes, http.StatusUnauthorized) - return - } - - w.WriteHeader(http.StatusOK) - if _, err := w.Write([]byte(`{"success": true}`)); err != nil { - require.NoError(t, err) - } - responseCodes = append(responseCodes, http.StatusOK) - })) - defer server.Close() - - // Create mock refresher with token sequence: first invalid, then valid - mockRef := &mockRefresher{ - token: &oauth2.Token{ - AccessToken: "invalid", - TokenType: "Bearer", - Expiry: time.Now().Add(time.Hour), - }, - } - - // Create our Transport with the mock refresher - transport := &Transport{ - base: &oauth2.Transport{ - Source: mockRef, - Base: http.DefaultTransport, - }, - refresher: mockRef, - } - - // Create HTTP client with our transport - client := &http.Client{Transport: transport} - - // Make a POST request with body content - requestBody := "test-body-content" - req, err := http.NewRequest("POST", server.URL, bytes.NewBufferString(requestBody)) - require.NoError(t, err) - req.Header.Set("Content-Type", "application/json") - - // Execute the request - resp, err := client.Do(req) - require.NoError(t, err) - defer resp.Body.Close() - - // Verify we got a successful response after retry - assert.Equal(t, http.StatusOK, resp.StatusCode) - - responseBody, err := io.ReadAll(resp.Body) - require.NoError(t, err) - assert.Equal(t, `{"success": true}`, string(responseBody)) - - // Assert the response codes (401 first and then a 200) - assert.Equal(t, []int{http.StatusUnauthorized, http.StatusOK}, responseCodes) -} diff --git a/internal/config/auth.go b/internal/config/auth.go new file mode 100644 index 000000000..7681b0ffa --- /dev/null +++ b/internal/config/auth.go @@ -0,0 +1,268 @@ +package config + +import ( + "errors" + "fmt" + "io/fs" + "os" + "path/filepath" + "regexp" + "strings" + + "gopkg.in/yaml.v3" +) + +// Auth holds the authentication settings. +// +// They are read from the same sources and with the same precedence as the legacy CLI: the embedded config, then the +// user's config.yaml file, then environment variables. +type Auth struct { + Token string // An API token (api.token). + TokenFile string // A file containing an API token (api.token_file), resolved to an absolute path. + AccessToken string // A raw access token (api.access_token), never stored or refreshed. + + SessionID string // The session ID (api.session_id), defaulting to "default". + SessionIDFromEnv bool // Whether the session ID was set via the SESSION_ID environment variable. + + BaseURL string // The API base URL (api.base_url). + AuthURL string // The auth server URL (api.auth_url). + AuthorizeURL string // The OAuth 2.0 authorization endpoint (api.oauth2_auth_url). + TokenURL string // The OAuth 2.0 token endpoint (api.oauth2_token_url). + RevokeURL string // The OAuth 2.0 revocation endpoint (api.oauth2_revoke_url). + ClientID string // The OAuth 2.0 client ID (api.oauth2_client_id). + + DisableCredentialHelpers bool // Whether to store credentials in files instead of the keychain. + SkipSSL bool // Whether to skip TLS verification (api.skip_ssl). + DisableLocks bool // Whether to skip locking (api.disable_locks). +} + +// authKey describes an auth config key under "api", and the environment variables (without the prefix) that set it. +type authKey struct { + key string + envVars []string // In ascending order of precedence. + target func(a *authSources) *string +} + +// authSources holds the raw string values, as they are overridden by each source. +type authSources struct { + token, tokenFile, accessToken, sessionID string + baseURL, authURL, authorizeURL, tokenURL, revokeURL, clientID string + disableCredentialHelpers, skipSSL, disableLocks string +} + +// authKeys lists the keys that are read, with their env var names. +// +// The generic form API_ comes first, then the legacy CLI's aliases, which take precedence. In the legacy CLI +// "API_TOKEN" is an alias for api.access_token (deprecated), overriding the generic form for api.token. +var authKeys = []authKey{ + {"token", []string{"TOKEN"}, func(a *authSources) *string { return &a.token }}, + {"token_file", []string{"API_TOKEN_FILE"}, func(a *authSources) *string { return &a.tokenFile }}, + {"access_token", []string{"API_ACCESS_TOKEN", "API_TOKEN"}, func(a *authSources) *string { return &a.accessToken }}, + {"session_id", []string{"API_SESSION_ID"}, func(a *authSources) *string { return &a.sessionID }}, + {"base_url", []string{"API_BASE_URL", "API_URL"}, func(a *authSources) *string { return &a.baseURL }}, + {"auth_url", []string{"API_AUTH_URL", "AUTH_URL"}, func(a *authSources) *string { return &a.authURL }}, + {"oauth2_auth_url", []string{"API_OAUTH2_AUTH_URL", "OAUTH2_AUTH_URL"}, + func(a *authSources) *string { return &a.authorizeURL }}, + {"oauth2_token_url", []string{"API_OAUTH2_TOKEN_URL", "OAUTH2_TOKEN_URL"}, + func(a *authSources) *string { return &a.tokenURL }}, + {"oauth2_revoke_url", []string{"API_OAUTH2_REVOKE_URL", "OAUTH2_REVOKE_URL"}, + func(a *authSources) *string { return &a.revokeURL }}, + {"oauth2_client_id", []string{"API_OAUTH2_CLIENT_ID", "OAUTH2_CLIENT_ID"}, + func(a *authSources) *string { return &a.clientID }}, + {"disable_credential_helpers", []string{"API_DISABLE_CREDENTIAL_HELPERS"}, + func(a *authSources) *string { return &a.disableCredentialHelpers }}, + {"skip_ssl", []string{"API_SKIP_SSL", "SKIP_SSL"}, func(a *authSources) *string { return &a.skipSSL }}, + {"disable_locks", []string{"API_DISABLE_LOCKS", "DISABLE_LOCKS"}, + func(a *authSources) *string { return &a.disableLocks }}, +} + +var sessionIDPattern = regexp.MustCompile(`(?i)^[a-z0-9_-]+$`) + +// ValidateSessionID checks a user-provided session ID. +func ValidateSessionID(id string) error { + if strings.HasPrefix(id, "api-token-") || !sessionIDPattern.MatchString(id) { + return fmt.Errorf("invalid session ID: %s", id) + } + return nil +} + +// Auth reads the authentication settings. +func (c *Config) Auth() (*Auth, error) { + src := &authSources{ + token: c.API.Token, + tokenFile: c.API.TokenFile, + accessToken: c.API.AccessToken, + baseURL: c.API.BaseURL, + authURL: c.API.AuthURL, + authorizeURL: c.API.OAuth2AuthorizeURL, + tokenURL: c.API.OAuth2TokenURL, + revokeURL: c.API.OAuth2RevokeURL, + clientID: c.API.OAuth2ClientID, + sessionID: c.API.SessionID, + } + if c.API.DisableCredentialHelpers { + src.disableCredentialHelpers = "1" + } + if c.API.SkipSSL { + src.skipSSL = "1" + } + if c.API.DisableLocks { + src.disableLocks = "1" + } + + userConfigDir, err := c.UserConfigDir() + if err != nil { + return nil, err + } + if err := applyUserAuthConfig(src, filepath.Join(userConfigDir, "config.yaml")); err != nil { + return nil, err + } + + prefix := c.Application.EnvPrefix + for _, k := range authKeys { + for _, v := range k.envVars { + if val, ok := os.LookupEnv(prefix + v); ok { + *k.target(src) = val + } + } + } + + a := &Auth{ + Token: src.token, + AccessToken: src.accessToken, + SessionID: src.sessionID, + BaseURL: src.baseURL, + AuthURL: src.authURL, + AuthorizeURL: src.authorizeURL, + TokenURL: src.tokenURL, + RevokeURL: src.revokeURL, + ClientID: src.clientID, + DisableCredentialHelpers: phpBool(src.disableCredentialHelpers), + SkipSSL: phpBool(src.skipSSL), + DisableLocks: phpBool(src.disableLocks), + } + + if src.tokenFile != "" { + a.TokenFile = src.tokenFile + if !filepath.IsAbs(a.TokenFile) && !strings.HasPrefix(a.TokenFile, `\`) { + a.TokenFile = filepath.Join(userConfigDir, a.TokenFile) + } + } + + // The session ID file is only read if SESSION_ID is not set. + if envID, ok := os.LookupEnv(prefix + "SESSION_ID"); ok { + a.SessionID = envID + } else { + fileID, err := c.readSessionIDFile() + if err != nil { + return nil, err + } + if fileID != "" { + a.SessionID = fileID + } + } + if a.SessionID == "" { + a.SessionID = "default" + } + if err := ValidateSessionID(a.SessionID); err != nil { + return nil, err + } + a.SessionIDFromEnv = a.SessionID != "default" && a.SessionID == os.Getenv(prefix+"SESSION_ID") + + if a.AuthURL != "" { + base := strings.TrimRight(a.AuthURL, "/") + for _, d := range []struct { + target *string + path string + }{ + {&a.AuthorizeURL, "/oauth2/authorize"}, + {&a.TokenURL, "/oauth2/token"}, + {&a.RevokeURL, "/oauth2/revoke"}, + } { + if *d.target == "" { + *d.target = base + d.path + } + } + } + if a.ClientID == "" { + a.ClientID = c.Application.Slug + } + + return a, nil +} + +// SessionIDFile returns the path to the file where the session ID is saved by session:switch. +func (c *Config) SessionIDFile() (string, error) { + dir, err := c.WritableUserDir() + if err != nil { + return "", err + } + return filepath.Join(dir, "session-id"), nil +} + +func (c *Config) readSessionIDFile() (string, error) { + path, err := c.SessionIDFile() + if err != nil { + return "", err + } + b, err := os.ReadFile(path) + if err != nil { + if errors.Is(err, fs.ErrNotExist) { + return "", nil + } + return "", err + } + id := strings.TrimSpace(string(b)) + if err := ValidateSessionID(id); err != nil { + return "", fmt.Errorf("invalid session ID in file: %s", path) + } + return id, nil +} + +// UserConfigDir returns the absolute path to the user config directory, e.g. ~/.upsun-cli. +func (c *Config) UserConfigDir() (string, error) { + home, err := c.HomeDir() + if err != nil { + return "", err + } + return filepath.Join(home, c.Application.UserConfigDir), nil +} + +// applyUserAuthConfig reads the "api" keys from the user's config file, if it exists. +func applyUserAuthConfig(src *authSources, path string) error { + b, err := os.ReadFile(path) + if err != nil { + if errors.Is(err, fs.ErrNotExist) { + return nil + } + return err + } + var userConfig struct { + API map[string]any `yaml:"api"` + } + if err := yaml.Unmarshal(b, &userConfig); err != nil { + return fmt.Errorf("invalid config file %s: %w", path, err) + } + for _, k := range authKeys { + v, ok := userConfig.API[k.key] + if !ok || v == nil { + continue + } + switch v := v.(type) { + case bool: + if v { + *k.target(src) = "1" + } else { + *k.target(src) = "" + } + default: + *k.target(src) = fmt.Sprint(v) + } + } + return nil +} + +// phpBool converts a string to a boolean in the same way as PHP. +func phpBool(s string) bool { + return s != "" && s != "0" +} diff --git a/internal/config/auth_test.go b/internal/config/auth_test.go new file mode 100644 index 000000000..3c0911bd2 --- /dev/null +++ b/internal/config/auth_test.go @@ -0,0 +1,204 @@ +package config_test + +import ( + "os" + "path/filepath" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/upsun/cli/internal/config" +) + +func TestAuth(t *testing.T) { + cases := []struct { + name string + env map[string]string + baseConfig string // Extra YAML for the "api" section of the base config. + userConfig string + sessionID string // The content of the session-id file. + check func(t *testing.T, a *config.Auth) + wantErr string + }{ + { + name: "defaults", + check: func(t *testing.T, a *config.Auth) { + assert.Equal(t, "default", a.SessionID) + assert.False(t, a.SessionIDFromEnv) + assert.Equal(t, "https://api.example.com", a.BaseURL) + assert.Equal(t, "https://auth.example.com/oauth2/authorize", a.AuthorizeURL) + assert.Equal(t, "https://auth.example.com/oauth2/token", a.TokenURL) + assert.Equal(t, "https://auth.example.com/oauth2/revoke", a.RevokeURL) + assert.Equal(t, "example-cli", a.ClientID) + assert.Empty(t, a.Token) + assert.False(t, a.DisableCredentialHelpers) + }, + }, + { + name: "env aliases", + env: map[string]string{ + "TOKEN": "api-token", + "API_TOKEN": "access-token", + "AUTH_URL": "https://auth2.example.com/", + "API_URL": "https://api2.example.com", + "API_BASE_URL": "https://ignored.example.com", + "SESSION_ID": "foo", + "SKIP_SSL": "1", + "DISABLE_LOCKS": "true", + }, + check: func(t *testing.T, a *config.Auth) { + assert.Equal(t, "api-token", a.Token) + assert.Equal(t, "access-token", a.AccessToken) + assert.Equal(t, "https://auth2.example.com/oauth2/token", a.TokenURL) + assert.Equal(t, "https://api2.example.com", a.BaseURL) + assert.Equal(t, "foo", a.SessionID) + assert.True(t, a.SessionIDFromEnv) + assert.True(t, a.SkipSSL) + assert.True(t, a.DisableLocks) + }, + }, + { + name: "env generic forms", + env: map[string]string{ //nolint:gosec // test values + "API_TOKEN_FILE": "token.txt", + "API_ACCESS_TOKEN": "access-token", + "API_AUTH_URL": "https://auth3.example.com", + "API_OAUTH2_TOKEN_URL": "https://token.example.com", + "API_OAUTH2_CLIENT_ID": "my-client", + "API_SESSION_ID": "bar", + "API_DISABLE_CREDENTIAL_HELPERS": "1", + "API_SKIP_SSL": "0", + }, + check: func(t *testing.T, a *config.Auth) { + assert.True(t, filepath.IsAbs(a.TokenFile)) + assert.Equal(t, "token.txt", filepath.Base(a.TokenFile)) + assert.Equal(t, "access-token", a.AccessToken) + assert.Equal(t, "https://auth3.example.com/oauth2/authorize", a.AuthorizeURL) + assert.Equal(t, "https://token.example.com", a.TokenURL) + assert.Equal(t, "my-client", a.ClientID) + assert.Equal(t, "bar", a.SessionID) + assert.False(t, a.SessionIDFromEnv) + assert.True(t, a.DisableCredentialHelpers) + assert.False(t, a.SkipSSL) + }, + }, + { + name: "empty env var overrides config", + env: map[string]string{"TOKEN": ""}, + userConfig: `api: + token: from-file +`, + check: func(t *testing.T, a *config.Auth) { + assert.Empty(t, a.Token) + }, + }, + { + name: "user config file", + userConfig: `api: + token: from-file + token_file: /abs/token + auth_url: https://auth4.example.com + disable_credential_helpers: true + session_id: baz +`, + check: func(t *testing.T, a *config.Auth) { + assert.Equal(t, "from-file", a.Token) + assert.Equal(t, "/abs/token", a.TokenFile) + assert.Equal(t, "https://auth4.example.com/oauth2/revoke", a.RevokeURL) + assert.True(t, a.DisableCredentialHelpers) + assert.Equal(t, "baz", a.SessionID) + }, + }, + { + name: "base config", + baseConfig: ` + token: from-base + token_file: /abs/base-token + access_token: base-access-token + disable_locks: true +`, + check: func(t *testing.T, a *config.Auth) { + assert.Equal(t, "from-base", a.Token) + assert.Equal(t, "/abs/base-token", a.TokenFile) + assert.Equal(t, "base-access-token", a.AccessToken) + assert.True(t, a.DisableLocks) + }, + }, + { + name: "user config overrides base config", + baseConfig: "\n token: from-base\n", + userConfig: "api: {token: from-file}\n", + check: func(t *testing.T, a *config.Auth) { + assert.Equal(t, "from-file", a.Token) + }, + }, + { + name: "env overrides user config", + env: map[string]string{"TOKEN": "from-env"}, + userConfig: "api: {token: from-file}\n", + check: func(t *testing.T, a *config.Auth) { + assert.Equal(t, "from-env", a.Token) + }, + }, + { + name: "session ID file", + sessionID: "from-file\n", + env: map[string]string{"API_SESSION_ID": "generic"}, + check: func(t *testing.T, a *config.Auth) { + assert.Equal(t, "from-file", a.SessionID) + }, + }, + { + name: "session ID env overrides file", + sessionID: "from-file", + env: map[string]string{"SESSION_ID": "from-env"}, + check: func(t *testing.T, a *config.Auth) { + assert.Equal(t, "from-env", a.SessionID) + }, + }, + { + name: "invalid session ID", + env: map[string]string{"SESSION_ID": "a/b"}, + wantErr: "invalid session ID: a/b", + }, + { + name: "reserved session ID", + env: map[string]string{"SESSION_ID": "api-token-abc"}, + wantErr: "invalid session ID", + }, + { + name: "invalid session ID file", + sessionID: "a b", + wantErr: "invalid session ID in file", + }, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + cnf, err := config.FromYAML([]byte(strings.Replace(validConfig, "api:", "api:"+c.baseConfig, 1))) + require.NoError(t, err) + home := t.TempDir() + t.Setenv("EXAMPLE_CLI_HOME", home) + for k, v := range c.env { + t.Setenv("EXAMPLE_CLI_"+k, v) + } + dir := filepath.Join(home, ".example-cli") + require.NoError(t, os.MkdirAll(dir, 0o700)) + if c.userConfig != "" { + require.NoError(t, os.WriteFile(filepath.Join(dir, "config.yaml"), []byte(c.userConfig), 0o600)) + } + if c.sessionID != "" { + require.NoError(t, os.WriteFile(filepath.Join(dir, "session-id"), []byte(c.sessionID), 0o600)) + } + + a, err := cnf.Auth() + if c.wantErr != "" { + assert.ErrorContains(t, err, c.wantErr) + return + } + require.NoError(t, err) + c.check(t, a) + }) + } +} diff --git a/internal/config/dir.go b/internal/config/dir.go index f366cacfa..635d5168e 100644 --- a/internal/config/dir.go +++ b/internal/config/dir.go @@ -59,7 +59,10 @@ func (c *Config) TempDir() (string, error) { return path, nil } -// WritableUserDir returns the path to a writable user-level directory. +// WritableUserDir returns the path to a writable user-level directory, e.g. for credentials and state. +// +// As in the legacy CLI, which shares it, a temporary directory is used if the directory in the home directory cannot +// be written, e.g. on an application container. // // Deprecated: unless backwards compatibility is desired, TempDir is preferable. func (c *Config) WritableUserDir() (string, error) { @@ -71,6 +74,9 @@ func (c *Config) WritableUserDir() (string, error) { return "", err } path := filepath.Join(hd, c.Application.WritableUserDir) + if !canWrite(path) { + path = filepath.Join(os.TempDir(), c.Application.TempSubDir) + } if err := os.MkdirAll(path, 0o700); err != nil { return "", err } @@ -79,10 +85,32 @@ func (c *Config) WritableUserDir() (string, error) { return path, nil } -// HomeDir returns the home directory configured via an environment variable, or the OS's user home directory otherwise. +// canWrite checks whether a directory is writable, or can be created, using permissions only. +// +// This matches the legacy CLI (Filesystem::canWrite), so both choose the same directory, e.g. even on a full disk. +func canWrite(path string) bool { + if info, err := os.Stat(path); err == nil { + return info.IsDir() && isWritable(path, info) + } + for p := filepath.Dir(path); ; p = filepath.Dir(p) { + if info, err := os.Stat(p); err == nil { + return isWritable(p, info) + } + if filepath.Dir(p) == p { + return false + } + } +} + +// HomeDir returns the user's home directory. +// +// It checks the same environment variables as the legacy CLI, in order: {ENV_PREFIX}HOME, HOME and USERPROFILE. +// On Windows, HOME can differ from USERPROFILE, e.g. in MSYS2 or Cygwin. func (c *Config) HomeDir() (string, error) { - if fromEnv := os.Getenv(c.Application.EnvPrefix + "HOME"); fromEnv != "" { - return fromEnv, nil + for _, name := range []string{c.Application.EnvPrefix + "HOME", "HOME", "USERPROFILE"} { + if v := os.Getenv(name); v != "" { + return v, nil + } } return os.UserHomeDir() } diff --git a/internal/config/dir_test.go b/internal/config/dir_test.go new file mode 100644 index 000000000..8e4b8a848 --- /dev/null +++ b/internal/config/dir_test.go @@ -0,0 +1,78 @@ +package config_test + +import ( + "os" + "path/filepath" + "runtime" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/upsun/cli/internal/config" +) + +func TestHomeDir(t *testing.T) { + cases := []struct { + name string + env map[string]string + want string + }{ + {"prefixed var first", map[string]string{"EXAMPLE_CLI_HOME": "/a", "HOME": "/b", "USERPROFILE": "/c"}, "/a"}, + {"then HOME", map[string]string{"EXAMPLE_CLI_HOME": "", "HOME": "/b", "USERPROFILE": "/c"}, "/b"}, + {"then USERPROFILE", map[string]string{"EXAMPLE_CLI_HOME": "", "HOME": "", "USERPROFILE": "/c"}, "/c"}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + cnf, err := config.FromYAML([]byte(validConfig)) + require.NoError(t, err) + for k, v := range c.env { + t.Setenv(k, v) + } + home, err := cnf.HomeDir() + require.NoError(t, err) + assert.Equal(t, c.want, home) + }) + } +} + +// TestWritableUserDir_TempFallback checks the cases where the legacy CLI uses a temporary directory instead. +func TestWritableUserDir_TempFallback(t *testing.T) { + cases := []struct { + name string + setup func(t *testing.T, home string) + }{ + { + name: "read-only home", + setup: func(t *testing.T, home string) { + if runtime.GOOS == "windows" || os.Geteuid() == 0 { + t.Skip("needs Unix permissions") + } + require.NoError(t, os.Chmod(home, 0o500)) + t.Cleanup(func() { _ = os.Chmod(home, 0o700) }) + }, + }, + { + name: "a file in place of the directory", + setup: func(t *testing.T, home string) { + require.NoError(t, os.WriteFile(filepath.Join(home, ".example-cli"), nil, 0o600)) + }, + }, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + cnf, err := config.FromYAML([]byte(validConfig)) + require.NoError(t, err) + home := t.TempDir() + c.setup(t, home) + tmp := t.TempDir() + t.Setenv("EXAMPLE_CLI_HOME", home) + t.Setenv("TMPDIR", tmp) // Unix + t.Setenv("TMP", tmp) // Windows + + dir, err := cnf.WritableUserDir() + require.NoError(t, err) + assert.Equal(t, filepath.Join(tmp, "example-cli-tmp"), dir) + }) + } +} diff --git a/internal/config/dir_unix.go b/internal/config/dir_unix.go new file mode 100644 index 000000000..9e5e65c17 --- /dev/null +++ b/internal/config/dir_unix.go @@ -0,0 +1,14 @@ +//go:build !windows + +package config + +import ( + "os" + + "golang.org/x/sys/unix" +) + +// isWritable checks write permission, like PHP's is_writable. +func isWritable(path string, _ os.FileInfo) bool { + return unix.Access(path, unix.W_OK) == nil +} diff --git a/internal/config/dir_windows.go b/internal/config/dir_windows.go new file mode 100644 index 000000000..67fc17d44 --- /dev/null +++ b/internal/config/dir_windows.go @@ -0,0 +1,8 @@ +package config + +import "os" + +// isWritable checks the read-only attribute, like PHP's is_writable on Windows. +func isWritable(_ string, info os.FileInfo) bool { + return info.Mode().Perm()&0o200 != 0 +} diff --git a/internal/config/schema.go b/internal/config/schema.go index 4824b54f3..b950ea0fa 100644 --- a/internal/config/schema.go +++ b/internal/config/schema.go @@ -47,13 +47,21 @@ type Config struct { CheckInterval int `validate:"omitempty" yaml:"check_interval,omitempty"` // seconds, defaults to 3600 } `validate:"omitempty"` - // Fields only needed by the PHP (legacy) CLI, at least for now. + // API and authentication settings. Go reads the auth keys via Config.Auth, which also applies overrides. API struct { BaseURL string `validate:"required,url" yaml:"base_url"` // e.g. "https://api.upsun.com" AuthURL string `validate:"omitempty,url" yaml:"auth_url,omitempty"` // e.g. "https://auth.upsun.com" - UserAgent string `validate:"omitempty" yaml:"user_agent,omitempty"` // a template - see UserAgent method - SessionID string `validate:"omitempty,ascii" yaml:"session_id,omitempty"` // the ID for the authentication session - defaults to "default" + UserAgent string `validate:"omitempty" yaml:"user_agent,omitempty"` // a template - see UserAgent method + SessionID string `validate:"omitempty,session_id" yaml:"session_id,omitempty"` // the ID for the authentication session - defaults to "default" + + Token string `validate:"omitempty" yaml:"token,omitempty"` // an API token + TokenFile string `validate:"omitempty" yaml:"token_file,omitempty"` // a file containing an API token + AccessToken string `validate:"omitempty" yaml:"access_token,omitempty"` // a raw access token + + DisableCredentialHelpers bool `validate:"omitempty" yaml:"disable_credential_helpers,omitempty"` // store credentials in files, not the keychain + SkipSSL bool `validate:"omitempty" yaml:"skip_ssl,omitempty"` // skip TLS verification (not recommended) + DisableLocks bool `validate:"omitempty" yaml:"disable_locks,omitempty"` // skip locking OAuth2ClientID string `validate:"omitempty" yaml:"oauth2_client_id,omitempty"` // e.g. "upsun-cli" OAuth2AuthorizeURL string `validate:"required_without=AuthURL,omitempty,url" yaml:"oauth2_auth_url,omitempty"` // e.g. "https://auth.upsun.com/oauth2/authorize" @@ -74,6 +82,11 @@ type Config struct { ConsoleURL string `validate:"omitempty,url" yaml:"console_url,omitempty"` // e.g. "https://console.upsun.com" DocsURL string `validate:"omitempty,url" yaml:"docs_url,omitempty"` // e.g. "https://docs.upsun.com" } `validate:"required"` + // BrowserLogin customizes the page shown by the local server during browser login. + BrowserLogin struct { + Body string `validate:"omitempty" yaml:"body,omitempty"` // HTML, with {{title}} and {{content}} placeholders + CSS string `validate:"omitempty" yaml:"css,omitempty"` + } `validate:"omitempty" yaml:"browser_login,omitempty"` SSH struct { DomainWildcards []string `validate:"required" yaml:"domain_wildcards"` // e.g. ["*.platform.sh"] } `validate:"required"` diff --git a/internal/config/validator.go b/internal/config/validator.go index 12a8326bf..d2c49e55e 100644 --- a/internal/config/validator.go +++ b/internal/config/validator.go @@ -22,4 +22,7 @@ func initCustomValidators(v *validator.Validate) { _ = v.RegisterValidation("version", func(fl validator.FieldLevel) bool { return fl.Field().Kind() == reflect.String && version.Validate(fl.Field().String()) }) + _ = v.RegisterValidation("session_id", func(fl validator.FieldLevel) bool { + return fl.Field().Kind() == reflect.String && ValidateSessionID(fl.Field().String()) == nil + }) } diff --git a/internal/legacy/legacy.go b/internal/legacy/legacy.go index 1e01fb91c..c4856eae2 100644 --- a/internal/legacy/legacy.go +++ b/internal/legacy/legacy.go @@ -38,6 +38,8 @@ type CLIWrapper struct { DisableInteraction bool ForceColor bool DebugLogFunc func(string, ...any) + // ExtraEnv is added to the command's environment. + ExtraEnv []string initOnce sync.Once _cacheDir string @@ -146,6 +148,10 @@ func (c *CLIWrapper) Exec(ctx context.Context, args ...string) error { envPrefix+"WRAPPED=1", envPrefix+"APPLICATION_VERSION="+c.Version, ) + // The legacy CLI gets tokens and auth state by running this executable's hidden auth:internal command. + if exe, err := os.Executable(); err == nil { + cmd.Env = append(cmd.Env, envPrefix+"WRAPPER_EXECUTABLE="+exe) + } if c.DisableInteraction { cmd.Env = append(cmd.Env, envPrefix+"NO_INTERACTION=1") } @@ -158,6 +164,7 @@ func (c *CLIWrapper) Exec(ctx context.Context, args ...string) error { c.Version, PHPVersion, )) + cmd.Env = append(cmd.Env, c.ExtraEnv...) if err := cmd.Run(); err != nil { return fmt.Errorf("could not run PHP CLI command: %w", err) } diff --git a/legacy/config-defaults.yaml b/legacy/config-defaults.yaml index c4ce1e7a7..4b77fdd7c 100644 --- a/legacy/config-defaults.yaml +++ b/legacy/config-defaults.yaml @@ -413,22 +413,3 @@ experimental: # Enable all experiments. all_experiments: false - -# Optional styling for browser login web page. -# The "css" string will be added as a stylesheet. -# The "body" string will replace the whole body content. -# Two template parameters are available: -# - {{title}} is replaced with a plain-text page title -# - {{content}} is replaced with the rest of the content HTML -browser_login: - # css: '' - body: | - - -

{{title}}

- - {{content}} diff --git a/legacy/phpstan-baseline.neon b/legacy/phpstan-baseline.neon index d465fa818..71cbe7110 100644 --- a/legacy/phpstan-baseline.neon +++ b/legacy/phpstan-baseline.neon @@ -1,95 +1,5 @@ parameters: ignoreErrors: - - - message: '#^Binary operation "\." between ''http\://127\.0\.0\.1\:'' and mixed results in an error\.$#' - identifier: binaryOp.invalid - count: 1 - path: resources/oauth-listener/index.php - - - - message: '#^Parameter \#1 \$message of method Platformsh\\Cli\\OAuth\\Listener\:\:reportError\(\) expects string\|null, mixed given\.$#' - identifier: argument.type - count: 1 - path: resources/oauth-listener/index.php - - - - message: '#^Parameter \#1 \(mixed\) of echo cannot be converted to string\.$#' - identifier: echo.nonString - count: 1 - path: resources/oauth-listener/index.php - - - - message: '#^Parameter \#2 \$error of method Platformsh\\Cli\\OAuth\\Listener\:\:reportError\(\) expects string\|null, mixed given\.$#' - identifier: argument.type - count: 1 - path: resources/oauth-listener/index.php - - - - message: '#^Parameter \#3 \$hint of method Platformsh\\Cli\\OAuth\\Listener\:\:reportError\(\) expects string\|null, mixed given\.$#' - identifier: argument.type - count: 1 - path: resources/oauth-listener/index.php - - - - message: '#^Parameter \#3 \$subject of function preg_replace_callback expects array\\|string, mixed given\.$#' - identifier: argument.type - count: 1 - path: resources/oauth-listener/index.php - - - - message: '#^Property Platformsh\\Cli\\OAuth\\Listener\:\:\$authMethods \(string\) does not accept mixed\.$#' - identifier: assign.propertyType - count: 1 - path: resources/oauth-listener/index.php - - - - message: '#^Property Platformsh\\Cli\\OAuth\\Listener\:\:\$authUrl \(string\) does not accept mixed\.$#' - identifier: assign.propertyType - count: 1 - path: resources/oauth-listener/index.php - - - - message: '#^Property Platformsh\\Cli\\OAuth\\Listener\:\:\$clientId \(string\) does not accept mixed\.$#' - identifier: assign.propertyType - count: 1 - path: resources/oauth-listener/index.php - - - - message: '#^Property Platformsh\\Cli\\OAuth\\Listener\:\:\$codeChallenge \(string\) does not accept mixed\.$#' - identifier: assign.propertyType - count: 1 - path: resources/oauth-listener/index.php - - - - message: '#^Property Platformsh\\Cli\\OAuth\\Listener\:\:\$file \(string\) does not accept mixed\.$#' - identifier: assign.propertyType - count: 1 - path: resources/oauth-listener/index.php - - - - message: '#^Property Platformsh\\Cli\\OAuth\\Listener\:\:\$maxAge \(string\|null\) does not accept mixed\.$#' - identifier: assign.propertyType - count: 1 - path: resources/oauth-listener/index.php - - - - message: '#^Property Platformsh\\Cli\\OAuth\\Listener\:\:\$prompt \(string\) does not accept mixed\.$#' - identifier: assign.propertyType - count: 1 - path: resources/oauth-listener/index.php - - - - message: '#^Property Platformsh\\Cli\\OAuth\\Listener\:\:\$scope \(string\) does not accept mixed\.$#' - identifier: assign.propertyType - count: 1 - path: resources/oauth-listener/index.php - - - - message: '#^Property Platformsh\\Cli\\OAuth\\Listener\:\:\$state \(string\) does not accept mixed\.$#' - identifier: assign.propertyType - count: 1 - path: resources/oauth-listener/index.php - - message: '#^Argument of an invalid type mixed supplied for foreach, only iterables are supported\.$#' identifier: foreach.nonIterable @@ -390,120 +300,6 @@ parameters: count: 1 path: src/Command/App/AppListCommand.php - - - message: '#^Call to an undefined method Platformsh\\Client\\Connection\\ConnectorInterface\:\:getOAuth2Provider\(\)\.$#' - identifier: method.notFound - count: 1 - path: src/Command/Auth/ApiTokenLoginCommand.php - - - - message: '#^Call to an undefined method Platformsh\\Client\\Connection\\ConnectorInterface\:\:saveToken\(\)\.$#' - identifier: method.notFound - count: 1 - path: src/Command/Auth/ApiTokenLoginCommand.php - - - - message: '#^Cannot call method getAccessToken\(\) on mixed\.$#' - identifier: method.nonObject - count: 1 - path: src/Command/Auth/ApiTokenLoginCommand.php - - - - message: '#^Cannot cast mixed to string\.$#' - identifier: cast.string - count: 2 - path: src/Command/Auth/ApiTokenLoginCommand.php - - - - message: '#^Parameter \#1 \$validator of method Symfony\\Component\\Console\\Question\\Question\:\:setValidator\(\) expects \(callable\(mixed\)\: mixed\)\|null, Closure\(string\|null\)\: non\-empty\-string given\.$#' - identifier: argument.type - count: 1 - path: src/Command/Auth/ApiTokenLoginCommand.php - - - - message: '#^Parameter \#2 \$accessToken of method Platformsh\\Cli\\Command\\Auth\\ApiTokenLoginCommand\:\:saveTokens\(\) expects League\\OAuth2\\Client\\Token\\AccessToken, mixed given\.$#' - identifier: argument.type - count: 1 - path: src/Command/Auth/ApiTokenLoginCommand.php - - - - message: '#^Binary operation "\." between '' Description\: '' and mixed results in an error\.$#' - identifier: binaryOp.invalid - count: 1 - path: src/Command/Auth/BrowserLoginCommand.php - - - - message: '#^Binary operation "\." between '' Hint\: '' and mixed results in an error\.$#' - identifier: binaryOp.invalid - count: 1 - path: src/Command/Auth/BrowserLoginCommand.php - - - - message: '#^Binary operation "\." between '' OAuth 2\.0 error\: …'' and mixed results in an error\.$#' - identifier: binaryOp.invalid - count: 1 - path: src/Command/Auth/BrowserLoginCommand.php - - - - message: '#^Cannot access offset ''code'' on mixed\.$#' - identifier: offsetAccess.nonOffsetAccessible - count: 1 - path: src/Command/Auth/BrowserLoginCommand.php - - - - message: '#^Cannot access offset ''error'' on mixed\.$#' - identifier: offsetAccess.nonOffsetAccessible - count: 1 - path: src/Command/Auth/BrowserLoginCommand.php - - - - message: '#^Cannot access offset ''error_description'' on mixed\.$#' - identifier: offsetAccess.nonOffsetAccessible - count: 2 - path: src/Command/Auth/BrowserLoginCommand.php - - - - message: '#^Cannot access offset ''error_hint'' on mixed\.$#' - identifier: offsetAccess.nonOffsetAccessible - count: 1 - path: src/Command/Auth/BrowserLoginCommand.php - - - - message: '#^Cannot access offset ''redirect_uri'' on mixed\.$#' - identifier: offsetAccess.nonOffsetAccessible - count: 1 - path: src/Command/Auth/BrowserLoginCommand.php - - - - message: '#^Method Platformsh\\Cli\\Command\\Auth\\BrowserLoginCommand\:\:getAccessToken\(\) should return array\ but returns array\.$#' - identifier: return.type - count: 1 - path: src/Command/Auth/BrowserLoginCommand.php - - - - message: '#^Parameter \#1 \$authCode of method Platformsh\\Cli\\Command\\Auth\\BrowserLoginCommand\:\:getAccessToken\(\) expects string, mixed given\.$#' - identifier: argument.type - count: 1 - path: src/Command/Auth/BrowserLoginCommand.php - - - - message: '#^Parameter \#1 \$env of method Symfony\\Component\\Process\\Process\:\:setEnv\(\) expects array\, array\\|bool\|float\|int\|string\|null\> given\.$#' - identifier: argument.type - count: 1 - path: src/Command/Auth/BrowserLoginCommand.php - - - - message: '#^Parameter \#1 \$messages of method Symfony\\Component\\Console\\Output\\OutputInterface\:\:writeln\(\) expects iterable\|string, mixed given\.$#' - identifier: argument.type - count: 1 - path: src/Command/Auth/BrowserLoginCommand.php - - - - message: '#^Parameter \#3 \$redirectUri of method Platformsh\\Cli\\Command\\Auth\\BrowserLoginCommand\:\:getAccessToken\(\) expects string, mixed given\.$#' - identifier: argument.type - count: 1 - path: src/Command/Auth/BrowserLoginCommand.php - - message: '#^Binary operation "\." between ''A verification code…'' and mixed results in an error\.$#' identifier: binaryOp.invalid @@ -5598,12 +5394,6 @@ parameters: count: 1 path: src/Service/Api.php - - - message: '#^Call to function is_array\(\) with array\ will always evaluate to true\.$#' - identifier: function.alreadyNarrowedType - count: 1 - path: src/Service/Api.php - - message: '#^Cannot access offset ''_endpoint'' on mixed\.$#' identifier: offsetAccess.nonOffsetAccessible @@ -5709,7 +5499,7 @@ parameters: - message: '#^Cannot cast mixed to string\.$#' identifier: cast.string - count: 6 + count: 3 path: src/Service/Api.php - @@ -5736,12 +5526,6 @@ parameters: count: 1 path: src/Service/Api.php - - - message: '#^Method Platformsh\\Cli\\Service\\Api\:\:getAccessToken\(\) should return string but returns mixed\.$#' - identifier: return.type - count: 1 - path: src/Service/Api.php - - message: '#^Method Platformsh\\Cli\\Service\\Api\:\:getEnvironmentTasks\(\) should return array\\> but returns array\\.$#' identifier: return.type @@ -5808,12 +5592,6 @@ parameters: count: 1 path: src/Service/Api.php - - - message: '#^Parameter \#1 \$body of method Platformsh\\Cli\\Service\\Api\:\:isApiTokenInvalid\(\) expects array\, array\ given\.$#' - identifier: argument.type - count: 1 - path: src/Service/Api.php - - message: '#^Parameter \#1 \$data of class Platformsh\\Client\\Model\\Deployment\\EnvironmentDeployment constructor expects array, mixed given\.$#' identifier: argument.type @@ -5850,12 +5628,6 @@ parameters: count: 1 path: src/Service/Api.php - - - message: '#^Parameter \#1 \$data of method Platformsh\\Cli\\Service\\Api\:\:isSsoSessionExpired\(\) expects array\, array\ given\.$#' - identifier: argument.type - count: 1 - path: src/Service/Api.php - - message: '#^Parameter \#1 \$id of method Platformsh\\Cli\\Service\\Api\:\:getOrganizationById\(\) expects string, mixed given\.$#' identifier: argument.type @@ -5874,12 +5646,6 @@ parameters: count: 1 path: src/Service/Api.php - - - message: '#^Parameter \#1 \$storage of method Platformsh\\Client\\Session\\Session\:\:setStorage\(\) expects Platformsh\\Client\\Session\\Storage\\SessionStorageInterface, Platformsh\\Client\\Session\\Storage\\SessionStorageInterface\|null given\.$#' - identifier: argument.type - count: 1 - path: src/Service/Api.php - - message: '#^Parameter \#1 \$string of function strlen expects string, mixed given\.$#' identifier: argument.type @@ -7152,12 +6918,6 @@ parameters: count: 1 path: src/Util/YamlParser.php - - - message: '#^Cannot access offset ''access_token'' on mixed\.$#' - identifier: offsetAccess.nonOffsetAccessible - count: 1 - path: tests/Command/Auth/BrowserLoginCommandTest.php - - message: '#^Argument of an invalid type mixed supplied for foreach, only iterables are supported\.$#' identifier: foreach.nonIterable diff --git a/legacy/resources/oauth-listener/index.php b/legacy/resources/oauth-listener/index.php deleted file mode 100644 index b2f4294fa..000000000 --- a/legacy/resources/oauth-listener/index.php +++ /dev/null @@ -1,258 +0,0 @@ -state = $_ENV['CLI_OAUTH_STATE']; - $this->authUrl = $_ENV['CLI_OAUTH_AUTH_URL']; - $this->clientId = $_ENV['CLI_OAUTH_CLIENT_ID']; - $this->file = $_ENV['CLI_OAUTH_FILE']; - $this->prompt = $_ENV['CLI_OAUTH_PROMPT']; - $this->codeChallenge = $_ENV['CLI_OAUTH_CODE_CHALLENGE']; - $this->scope = $_ENV['CLI_OAUTH_SCOPE'] ?? ''; - $this->localUrl = 'http://127.0.0.1:' . $_SERVER['SERVER_PORT']; - $this->response = new Response(); - $this->authMethods = $_ENV['CLI_OAUTH_METHODS'] ?? ''; - $this->maxAge = $_ENV['CLI_OAUTH_MAX_AGE'] ?? null; - } - - /** - * @return string - */ - private function getOAuthUrl(): string - { - $params = [ - 'redirect_uri' => $this->localUrl, - 'state' => $this->state, - 'client_id' => $this->clientId, - 'prompt' => $this->prompt, - 'response_type' => 'code', - 'code_challenge' => $this->codeChallenge, - 'code_challenge_method' => 'S256', - 'scope' => $this->scope, - ]; - - if (!empty($this->authMethods)) { - $params['amr'] = $this->authMethods; - } - if ($this->maxAge !== null && $this->maxAge !== '') { - $params['max_age'] = $this->maxAge; - } - - return $this->authUrl . '?' . http_build_query($params, '', '&', PHP_QUERY_RFC3986); - } - - /** - * Check state, run logic, set page content. - */ - public function run(): void - { - // Respond after a successful OAuth2 redirect. - if (isset($_GET['state'], $_GET['code'])) { - if ($_GET['state'] !== $this->state) { - $this->reportError('Invalid state parameter'); - return; - } - if (isset($_GET['code_challenge']) && $_GET['code_challenge'] !== $this->codeChallenge) { - $this->reportError('Invalid returned code_challenge parameter'); - return; - } - if (!$this->sendToTerminal(['code' => $_GET['code'], 'redirect_uri' => $this->localUrl])) { - $this->reportError('Failed to send authorization code back to terminal'); - return; - } - $this->setRedirect($this->localUrl . '/?done'); - $this->response->content = '

Authentication response received, please wait...

'; - - return; - } - - // Show the final result page. - if (array_key_exists('done', $_GET)) { - $this->response->title = 'Successfully logged in'; - $this->response->content = '

You can return to the command line

'; - - return; - } - - // Respond after an OAuth2 error. - if (isset($_GET['error'])) { - $message = $_GET['error_description'] ?? null; - $hint = $_GET['error_hint'] ?? null; - $this->reportError($message, $_GET['error'], $hint); - return; - } - - // In any other case: redirect to login. - $url = $this->getOAuthUrl(); - $this->setRedirect($url); - $this->response->content = '

Log in.

'; - } - - /** - * @param string $url - * @param int $code - */ - private function setRedirect(string $url, int $code = 302): void - { - $this->response->code = $code; - $this->response->headers['Location'] = $url; - } - - public function getResponse(): Response - { - return $this->response; - } - - /** - * @param array $response - * - * @return bool - */ - private function sendToTerminal(array $response): bool - { - return (bool) file_put_contents($this->file, json_encode($response), LOCK_EX); - } - - /** - * @param string|null $message The error message. - * @param string|null $error An OAuth2 error type. - * @param string|null $hint An OAuth2 error hint. - */ - private function reportError(?string $message = null, ?string $error = null, ?string $hint = null): void - { - $this->response->headers['Status'] = '401'; - $this->response->title = 'Error'; - if (isset($error)) { - $this->response->content .= '

' . htmlspecialchars($error) . '

'; - } - if (isset($message)) { - $this->response->content .= '

' . htmlspecialchars($message) . '

'; - } - if (isset($hint)) { - $this->response->content .= '

' . htmlspecialchars($hint) . '

'; - } - if ($message || $error || $hint) { - $response = ['error' => $error, 'error_description' => $message, 'error_hint' => $hint]; - if (!$this->sendToTerminal($response)) { - $this->response->content .= '

Additionally: failed to send error message back to terminal

'; - } - } - $this->response->content .= '

Please try again

'; - } -} - -class Response -{ - /** @var array */ - public array $headers = []; - public int $code = 200; - public string $headTitle = ''; - public string $title = ''; - public string $content = ''; - - public function __construct() - { - // Set default title and headers. - $appName = getenv('CLI_OAUTH_APP_NAME') ?: 'CLI'; - $this->headTitle = htmlspecialchars($appName) . ': Authentication (temporary URL)'; - $this->headers = [ - 'Cache-Control' => 'no-cache', - 'Content-Type' => 'text/html; charset=utf-8', - ]; - } -} - -$configJson = file_get_contents('config.json'); -if ($configJson === false) { - throw new \RuntimeException('Failed to load configuration file: config.json'); -} -$config = (array) json_decode($configJson, true); - -$listener = new Listener(); -$listener->run(); - -$response = $listener->getResponse(); - -if (!empty($config['body'])) { - $body = preg_replace_callback('/\{\{\s*(content|title)\s*}}/', function (array $matches) use ($response) { - return ['content' => $response->content, 'title' => $response->title][$matches[1]]; - }, $config['body']); -} else { - $body = '

' . $response->title . '

' . $response->content; -} - -http_response_code($response->code); -foreach ($response->headers as $name => $value) { - header($name . ': ' . $value); -} -?> - - - - - <?php echo $response->headTitle; ?> - - - - - - - - - diff --git a/legacy/src/ApiToken/CredentialHelperStorage.php b/legacy/src/ApiToken/CredentialHelperStorage.php deleted file mode 100644 index eadb06fd0..000000000 --- a/legacy/src/ApiToken/CredentialHelperStorage.php +++ /dev/null @@ -1,49 +0,0 @@ -serverUrl = sprintf( - '%s/%s/api-token', - $config->getStr('application.slug'), - $config->getSessionId(), - ); - } - - /** - * @inheritDoc - */ - public function getToken(): string - { - return $this->manager->get($this->serverUrl) ?: ''; - } - - /** - * @inheritDoc - */ - public function storeToken($value): void - { - $this->manager->store($this->serverUrl, $value); - } - - /** - * @inheritDoc - */ - public function deleteToken(): void - { - $this->manager->erase($this->serverUrl); - } -} diff --git a/legacy/src/ApiToken/FileStorage.php b/legacy/src/ApiToken/FileStorage.php deleted file mode 100644 index 8d75edcc6..000000000 --- a/legacy/src/ApiToken/FileStorage.php +++ /dev/null @@ -1,101 +0,0 @@ -fs = $fs ?: new SymfonyFilesystem(); - } - - /** - * Loads the API token. - * - * @return string - */ - public function getToken(): string - { - return $this->load(); - } - - /** - * Stores an API token. - * - * @param string $value - */ - public function storeToken(string $value): void - { - $this->save($value); - } - - /** - * Deletes the saved token. - */ - public function deleteToken(): void - { - $this->save(''); - } - - private function save(string $token): void - { - $filename = $this->getFilename(); - if (empty($token)) { - if (file_exists($filename)) { - $this->fs->remove($filename); - } - return; - } - - // Avoid overwriting an already configured token file. - if (file_exists($filename) && $this->config->has('api.token_file') && $this->resolveTokenFile($this->config->getStr('api.token_file')) === $filename) { - throw new \RuntimeException('Failed to save API token: it would conflict with the existing api.token_file configuration.'); - } - - $this->fs->dumpFile($filename, $token); - $this->fs->chmod($filename, 0o600); - } - - /** - * @return string - */ - private function load(): string - { - $filename = $this->getFilename(); - if (file_exists($filename)) { - return trim((string) file_get_contents($filename)); - } - - return ''; - } - - /** - * @return string - */ - private function getFilename(): string - { - return $this->config->getSessionDir(true) . DIRECTORY_SEPARATOR . 'api-token'; - } - - /** - * Makes a relative path absolute, based on the user config dir. - */ - private function resolveTokenFile(string $tokenFile): string - { - if (!str_starts_with($tokenFile, '/') && !str_starts_with($tokenFile, '\\')) { - $tokenFile = $this->config->getUserConfigDir() . '/' . $tokenFile; - } - - return $tokenFile; - } -} diff --git a/legacy/src/ApiToken/Storage.php b/legacy/src/ApiToken/Storage.php deleted file mode 100644 index 7cdd9878b..000000000 --- a/legacy/src/ApiToken/Storage.php +++ /dev/null @@ -1,24 +0,0 @@ -isSupported()) { - return new CredentialHelperStorage($config, $manager); - } - - return new FileStorage($config); - } -} diff --git a/legacy/src/ApiToken/StorageInterface.php b/legacy/src/ApiToken/StorageInterface.php deleted file mode 100644 index e40671036..000000000 --- a/legacy/src/ApiToken/StorageInterface.php +++ /dev/null @@ -1,25 +0,0 @@ -config->getStr('service.name'); - $executable = $this->config->getStr('application.executable'); - - $help = 'Use this command to log in to your ' . $service . ' account using an API token.'; - if ($this->config->has('service.register_url')) { - $help .= "\n\nYou can create an account at:\n " . $this->config->getStr('service.register_url') . ''; - } - if ($this->config->has('service.api_tokens_url')) { - $help .= "\n\nIf you have an account, but you do not already have an API token, you can create one here:\n " - . $this->config->getStr('service.api_tokens_url') . ''; - } - $help .= "\n\nAlternatively, to log in to the CLI with a browser, run:\n " . $executable . ' auth:browser-login'; - $this->setHelp($help); - } - - protected function execute(InputInterface $input, OutputInterface $output): int - { - if ($this->api->hasApiToken(false)) { - $this->stdErr->writeln('An API token is already set via config'); - return 1; - } - if (!$input->isInteractive()) { - $this->stdErr->writeln('Non-interactive use of this command is not supported.'); - $this->stdErr->writeln("\n" . $this->login->getNonInteractiveAuthHelp('comment')); - return 1; - } - - $validator = function (?string $apiToken): string { - $apiToken = trim((string) $apiToken); - if (!strlen($apiToken)) { - throw new \RuntimeException('The token cannot be empty'); - } - - try { - $provider = $this->api->getClient(false)->getConnector()->getOAuth2Provider(); - $token = $provider->getAccessToken(new ApiToken(), [ - 'api_token' => $apiToken, - ]); - } catch (BadResponseException $e) { - if ($this->exceptionMeansInvalidToken($e)) { - throw new \RuntimeException('Invalid API token'); - } - throw $e; - } - - // Finalise login. - $this->stdErr->writeln(''); - $this->stdErr->writeln('The API token is valid.'); - $this->saveTokens($apiToken, $token); - - return $apiToken; - }; - $question = new Question("Please enter an API token:\n> "); - $question->setValidator($validator); - $question->setMaxAttempts(5); - $question->setHidden(true); - $this->questionHelper->ask($input, $output, $question); - - $this->login->finalize(); - - return 0; - } - - /** - * Saves the new tokens and safely logs out of the previous session. - * - * @param string $apiToken - * @param AccessToken $accessToken - */ - private function saveTokens(string $apiToken, AccessToken $accessToken): void - { - $this->api->logout(); - $this->tokenConfig->storage()->storeToken($apiToken); - - $this->api - ->getClient(false, true) - ->getConnector() - ->saveToken($accessToken); - } - - /** - * @param \Exception $e - * - * @return bool - */ - private function exceptionMeansInvalidToken(\Exception $e): bool - { - if (!$e instanceof BadResponseException || !in_array($e->getResponse()->getStatusCode(), [400, 401], true)) { - return false; - } - $json = (array) Utils::jsonDecode((string) $e->getResponse()->getBody(), true); - // Compatibility with legacy auth provider. - if (isset($json['error'], $json['error_description']) - && $json['error'] === 'invalid_grant' - && stripos((string) $json['error_description'], 'Invalid API token') !== false) { - return true; - } - // Compatibility with new auth provider. - if (isset($json['error'], $json['error_hint']) - && $json['error'] === 'request_unauthorized' - && stripos((string) $json['error_hint'], 'API token') !== false) { - return true; - } - - return false; - } -} diff --git a/legacy/src/Command/Auth/AuthTokenCommand.php b/legacy/src/Command/Auth/AuthTokenCommand.php deleted file mode 100644 index f1251d8ed..000000000 --- a/legacy/src/Command/Auth/AuthTokenCommand.php +++ /dev/null @@ -1,66 +0,0 @@ -addOption('header', 'H', InputOption::VALUE_NONE, 'Prefix the token with "' . self::RFC6750_PREFIX . '" to make an RFC 6750 header') - ->addOption('no-warn', 'W', InputOption::VALUE_NONE, 'Suppress the warning that is printed by default to stderr.' - . ' This option is preferred over redirecting stderr, as that would hide other potentially useful messages.'); - $help = \wordwrap( - 'This command prints a valid OAuth 2 access token to stdout. It can be used to make API requests via standard Bearer authentication (RFC 6750).' - . "\n\n" . 'Warning: access tokens must be kept secret.' - . "\n\n" . 'Using this command is not generally recommended, as it increases the chance of the token being leaked.' - . ' Take care not to expose the token in a shared program or system, or to send the token to the wrong API domain.', - ); - $executable = $this->config->getStr('application.executable'); - $apiUrl = $this->config->getApiUrl(); - $examples = [ - 'Print the payload for JWT-formatted tokens' => \sprintf('%s auth:token -W | cut -d. -f2 | base64 -d', $executable), - 'Use the token in a curl command' => \sprintf('curl -H"$(%s auth:token -HW)" %s/users/me', $executable, rtrim($apiUrl, '/')), - ]; - $help .= "\n\nExamples:"; - foreach ($examples as $description => $example) { - $help .= "\n\n$description:\n $example"; - } - $this->setHelp($help); - } - - protected function execute(InputInterface $input, OutputInterface $output): int - { - if (!Option::bool($input, 'no-warn')) { - $this->stdErr->writeln( - 'Warning: keep access tokens secret.', - ); - } - - $token = $this->api->getAccessToken(); - - $output->write(Option::bool($input, 'header') ? self::RFC6750_PREFIX . $token : $token); - - return 0; - } -} diff --git a/legacy/src/Command/Auth/BrowserLoginCommand.php b/legacy/src/Command/Auth/BrowserLoginCommand.php deleted file mode 100644 index c7b609607..000000000 --- a/legacy/src/Command/Auth/BrowserLoginCommand.php +++ /dev/null @@ -1,373 +0,0 @@ -config->getStr('application.name'); - - $this - ->addOption('force', 'f', InputOption::VALUE_NONE, 'Log in again, even if already logged in') - ->addOption('method', null, InputOption::VALUE_REQUIRED | InputOption::VALUE_IS_ARRAY, 'Require specific authentication method(s)') - ->addOption('max-age', null, InputOption::VALUE_REQUIRED, 'The maximum age (in seconds) of the web authentication session'); - Url::configureInput($this->getDefinition()); - - $executable = $this->config->getStr('application.executable'); - $help = 'Use this command to log in to the ' . $applicationName . ' using a web browser.' - . "\n\nIt launches a temporary local website which redirects you to log in if necessary, and then captures the resulting authorization code." - . "\n\nYour system's default browser will be used. You can override this using the --browser option." - . "\n\nAlternatively, to log in using an API token (without a browser), run: $executable auth:api-token-login" - . "\n\n" . $this->login->getNonInteractiveAuthHelp(); - $this->setHelp(\wordwrap($help, 80)); - } - - protected function execute(InputInterface $input, OutputInterface $output): int - { - $maxAge = Option::intOrNull($input, 'max-age'); - if ($this->api->hasApiToken(false)) { - $this->stdErr->writeln('Cannot log in via the browser, because an API token is set via config.'); - return 1; - } - if (!$input->isInteractive()) { - $this->stdErr->writeln('Non-interactive use of this command is not supported.'); - $this->stdErr->writeln("\n" . $this->login->getNonInteractiveAuthHelp('comment')); - return 1; - } - if ($this->config->getSessionId() !== 'default' || count($this->api->listSessionIds()) > 1) { - $this->stdErr->writeln(sprintf('The current session ID is: %s', $this->config->getSessionId())); - if (!$this->config->isSessionIdFromEnv()) { - $this->stdErr->writeln(sprintf('Change this using: %s session:switch', $this->config->getStr('application.executable'))); - } - $this->stdErr->writeln(''); - } - $connector = $this->api->getClient(false)->getConnector(); - $force = Option::bool($input, 'force'); - if (!$force && Option::stringArray($input, 'method') === [] && $maxAge === null && $connector->isLoggedIn()) { - // Get account information, simultaneously checking whether the API - // login is still valid. If the request works, then do not log in - // again (unless --force is used). If the request fails, proceed - // with login. - $api = $this->api; - try { - $api->inLoginCheck = true; - - $account = $api->getMyAccount(); - $this->stdErr->writeln(\sprintf( - 'You are already logged in as %s (%s)', - $account['username'], - $account['email'], - )); - - if (!$this->questionHelper->confirm('Log in anyway?', false)) { - return 1; - } - $force = true; - } catch (BadResponseException $e) { - if (in_array($e->getResponse()->getStatusCode(), [400, 401], true)) { - $this->io->debug('Already logged in, but a test request failed. Continuing with login.'); - } else { - throw $e; - } - } finally { - $api->inLoginCheck = false; - } - } - - // Set up the local PHP web server, which will serve an OAuth2 redirect - // and wait for the response. - // Firstly, find an address. The port needs to be within a known range, - // for validation by the remote server. - try { - $start = 5000; - $end = 5010; - $port = PortUtil::getPort($start, null, $end); - } catch (\Exception $e) { - if (stripos($e->getMessage(), 'failed to find') !== false) { - $this->stdErr->writeln(sprintf('Failed to find an available port between %d and %d.', $start, $end)); - $this->stdErr->writeln('Check if you have unnecessary services running on these ports.'); - $this->stdErr->writeln(sprintf('For more options, run: %s help login', $this->config->getStr('application.executable'))); - - return 1; - } - throw $e; - } - $localAddress = '127.0.0.1:' . $port; - $localUrl = 'http://' . $localAddress; - - // Then create the document root for the local server. This needs to be - // outside the CLI itself (since the CLI may be run as a Phar). - $listenerDir = $this->config->getWritableUserDir() . '/oauth-listener'; - $this->createDocumentRoot($listenerDir); - - // Create the file where a response will be saved (by the local server - // script). - $responseFile = $listenerDir . '/.response'; - if (file_put_contents($responseFile, '', LOCK_EX) === false) { - throw new \RuntimeException('Failed to create temporary file: ' . $responseFile); - } - chmod($responseFile, 0o600); - - // Start the local server. - $phpCommand = [(new PhpExecutableFinder())->find() ?: PHP_BINARY]; - if (!php_ini_loaded_file() && !php_ini_scanned_files()) { - // This process parsed no php.ini file, so tell the server not to - // parse one either. Otherwise it may load a php.ini which does not - // suit this PHP build, for example one that enables extensions - // which are already built in. - $phpCommand[] = '-n'; - } - $phpCommand[] = '-dvariables_order=egps'; - $process = new Process(array_merge($phpCommand, ['-S', $localAddress, '-t', $listenerDir])); - $codeVerifier = $this->generateCodeVerifier(); - $process->setEnv([ - 'CLI_OAUTH_APP_NAME' => $this->config->getStr('application.name'), - 'CLI_OAUTH_STATE' => $this->generateCodeVerifier(), // the state can just be any random string - 'CLI_OAUTH_CODE_CHALLENGE' => $this->convertVerifierToChallenge($codeVerifier), - 'CLI_OAUTH_AUTH_URL' => $this->config->get('api.oauth2_auth_url'), - 'CLI_OAUTH_CLIENT_ID' => $this->config->get('api.oauth2_client_id'), - 'CLI_OAUTH_PROMPT' => $force ? 'consent select_account' : 'consent', - 'CLI_OAUTH_SCOPE' => 'offline_access', - 'CLI_OAUTH_FILE' => $responseFile, - 'CLI_OAUTH_METHODS' => implode(' ', ArrayArgument::getOption($input, 'method')), - 'CLI_OAUTH_MAX_AGE' => (string) $maxAge, - ] + getenv()); - $process->setTimeout(null); - $this->stdErr->writeln('Starting local web server with command: ' . $process->getCommandLine() . '', OutputInterface::VERBOSITY_VERY_VERBOSE); - $process->start(); - - // Give the local server some time to start before checking its status - // or opening the browser (0.5 seconds). - usleep(500000); - - // Check the local server status. - if (!$process->isRunning()) { - $this->stdErr->writeln('Failed to start local web server.'); - $this->stdErr->writeln(trim($process->getErrorOutput())); - - return 1; - } - if ($this->url->openUrl($localUrl, false)) { - $this->stdErr->writeln(sprintf('Opened URL: %s', $localUrl)); - $this->stdErr->writeln('Please use the browser to log in.'); - } else { - $this->stdErr->writeln('Please open the following URL in a browser and log in:'); - $this->stdErr->writeln('' . $localUrl . ''); - } - - // Show some help. - $this->stdErr->writeln(''); - $this->stdErr->writeln('Help:'); - $this->stdErr->writeln(' Leave this command running during login.'); - $this->stdErr->writeln(' If you need to quit, use Ctrl+C.'); - $this->stdErr->writeln(''); - - // Wait for the file to be filled with an OAuth2 authorization code. - /** @var null|array{code: string, redirect_uri: string}|array{error: string, error_description: string, error_hint: string} $response */ - $response = null; - $start = time(); - while ($process->isRunning()) { - usleep(300000); - if (!file_exists($responseFile)) { - $this->stdErr->writeln('File not found: ' . $responseFile . ''); - $this->stdErr->writeln(''); - break; - } - $responseRaw = file_get_contents($responseFile); - if ($responseRaw === false) { - $this->stdErr->writeln('Failed to read file: ' . $responseFile . ''); - $this->stdErr->writeln(''); - break; - } - if ($responseRaw !== '') { - $response = json_decode($responseRaw, true); - break; - } - if (time() - $start >= 1800) { - $this->stdErr->writeln('Login timed out after 30 minutes'); - $this->stdErr->writeln(''); - break; - } - } - - // Allow a little time for the final page to be displayed in the - // browser. - usleep(100000); - - // Clean up. - $process->stop(); - (new Filesystem())->remove([$listenerDir]); - - if (empty($response) || empty($response['code'])) { - $this->stdErr->writeln('Failed to get an authorization code.'); - $this->stdErr->writeln(''); - if (!empty($response['error']) && !empty($response['error_description'])) { - $this->stdErr->writeln(' OAuth 2.0 error: ' . $response['error'] . ''); - $this->stdErr->writeln(' Description: ' . $response['error_description']); - if (!empty($response['error_hint'])) { - $this->stdErr->writeln(' Hint: ' . $response['error_hint']); - } - $this->stdErr->writeln(''); - } elseif (!empty($response['error_description'])) { - $this->stdErr->writeln($response['error_description']); - $this->stdErr->writeln(''); - } - $this->stdErr->writeln('Please try again.'); - - return 1; - } - - $code = $response['code']; - - // Using the authorization code, request an access token. - $this->stdErr->writeln('Login information received. Verifying...'); - $token = $this->getAccessToken($code, $codeVerifier, $response['redirect_uri'] ?? $localUrl); - - // Finalize login: log out and save the new credentials. - $this->api->logout(); - - // Save the new tokens to the persistent session. - $session = $this->api->getClient(false)->getConnector()->getSession(); - $this->saveAccessToken($token, $session); - - $this->login->finalize(); - - if (empty($token['refresh_token'])) { - $this->stdErr->writeln(''); - $clientId = $this->config->getStr('api.oauth2_client_id'); - $this->stdErr->writeln([ - 'Warning:', - 'No refresh token is available. This will cause frequent login errors.', - 'Please contact support.', - "For internal use: the OAuth 2 client is probably misconfigured (client ID: $clientId).", - ]); - } - - return 0; - } - - /** - * @param array $tokenData - * @param SessionInterface $session - */ - private function saveAccessToken(array $tokenData, SessionInterface $session): void - { - $token = new AccessToken($tokenData); - $session->set('accessToken', $token->getToken()); - $session->set('tokenType', $tokenData['token_type'] ?: null); - $session->set('expires', $token->getExpires()); - $session->set('refreshToken', $token->getRefreshToken()); - $session->save(); - } - - /** - * @param string $dir - */ - private function createDocumentRoot(string $dir): void - { - if (!is_dir($dir) && !mkdir($dir, 0o700, true)) { - throw new \RuntimeException('Failed to create temporary directory: ' . $dir); - } - if (!file_put_contents($dir . '/index.php', (string) file_get_contents(CLI_ROOT . '/resources/oauth-listener/index.php'))) { - throw new \RuntimeException('Failed to write temporary file: ' . $dir . '/index.php'); - } - if (!file_put_contents($dir . '/config.json', (string) json_encode((array) $this->config->get('browser_login'), JSON_UNESCAPED_SLASHES))) { - throw new \RuntimeException('Failed to write temporary file: ' . $dir . '/config.json'); - } - } - - /** - * Exchanges the authorization code for an access token. - * - * @return array - */ - private function getAccessToken(string $authCode, string $codeVerifier, string $redirectUri): array - { - // Use the shared HTTP client, so that the request uses the detected CA - // bundle and any configured proxy. - $client = $this->api->getExternalHttpClient(); - $request = new Request('POST', $this->config->getStr('api.oauth2_token_url'), body: http_build_query([ - 'grant_type' => 'authorization_code', - 'code' => $authCode, - 'redirect_uri' => $redirectUri, - 'code_verifier' => $codeVerifier, - ])); - - try { - $response = $client->send($request, [ - 'headers' => [ - 'Content-Type' => 'application/x-www-form-urlencoded', - ], - 'auth' => [$this->config->get('api.oauth2_client_id'), ''], - ]); - - return (array) Utils::jsonDecode((string) $response->getBody(), true); - } catch (BadResponseException $e) { - throw ApiResponseException::create($request, $e->getResponse(), $e); - } - } - - /** - * Gets a PKCE code verifier to use with the OAuth2 code request. - */ - private function generateCodeVerifier(): string - { - // This uses paragonie/random_compat as a polyfill for PHP < 7.0. - return $this->base64UrlEncode(random_bytes(32)); - } - - /** - * Base64URL-encodes a string according to the PKCE spec. - * - * @see https://tools.ietf.org/html/rfc7636 - * - * @param string $data - * - * @return string - */ - private function base64UrlEncode(string $data): string - { - return str_replace(['+', '/'], ['-', '_'], rtrim(base64_encode($data), '=')); - } - - /** - * Generates a PKCE code challenge using the S256 transformation on a verifier. - */ - private function convertVerifierToChallenge(string $verifier): string - { - return $this->base64UrlEncode(hash('sha256', $verifier, true)); - } -} diff --git a/legacy/src/Command/Auth/ExportSessionsCommand.php b/legacy/src/Command/Auth/ExportSessionsCommand.php new file mode 100644 index 000000000..175396b9a --- /dev/null +++ b/legacy/src/Command/Auth/ExportSessionsCommand.php @@ -0,0 +1,163 @@ +addOption('delete', null, InputOption::VALUE_NONE, 'Delete the stored sessions, without revoking them'); + } + + protected function execute(InputInterface $input, OutputInterface $output): int + { + if (Option::bool($input, 'delete')) { + $this->delete(); + return 0; + } + + /** @var array> $sessions */ + $sessions = []; + // Values are only set once: the keychain takes precedence over files, as the legacy CLI only used files + // when the keychain was unavailable. + $set = function (string $id, string $key, mixed $value) use (&$sessions): void { + if ($value !== null && $value !== '' && $value !== false && !isset($sessions[$id][$key])) { + $sessions[$id][$key] = $value; + } + }; + $setFromSession = function (string $id, mixed $data) use ($set): void { + if (!is_array($data)) { + return; + } + $set($id, 'access_token', $data['accessToken'] ?? null); + $set($id, 'refresh_token', $data['refreshToken'] ?? null); + $set($id, 'token_type', $data['tokenType'] ?? null); + $set($id, 'expires', isset($data['expires']) && is_numeric($data['expires']) ? (int) $data['expires'] : null); + }; + + // Sessions and API tokens in the keychain, if the credential helper is already installed. + $manager = $this->credentialHelper(); + if ($manager !== null) { + $prefix = $this->config->getStr('application.slug') . '/'; + foreach (array_keys($manager->listAll()) as $url) { + if (!str_starts_with((string) $url, $prefix)) { + continue; + } + $path = substr((string) $url, strlen($prefix)); + $secret = $manager->get((string) $url); + if ($secret === false) { + continue; + } + if (str_ends_with($path, '/api-token')) { + $set(substr($path, 0, -strlen('/api-token')), 'api_token', trim($secret)); + } else { + $setFromSession($path, json_decode((string) base64_decode($secret, true), true)); + } + } + } + + // Sessions and API tokens in files. + foreach ($this->sessionFiles() as $id => $file) { + $setFromSession($id, json_decode((string) file_get_contents($file), true)); + } + foreach ($this->apiTokenFiles() as $id => $file) { + $set($id, 'api_token', trim((string) file_get_contents($file))); + } + + $output->writeln((string) json_encode((object) $sessions, JSON_UNESCAPED_SLASHES)); + + return 0; + } + + /** + * Deletes the exported copies. SSH certificates and the credential helper itself are kept. + */ + private function delete(): void + { + $manager = $this->credentialHelper(); + if ($manager !== null) { + $prefix = $this->config->getStr('application.slug') . '/'; + foreach (array_keys($manager->listAll()) as $url) { + if (str_starts_with((string) $url, $prefix)) { + $manager->erase((string) $url); + } + } + } + $fs = new Filesystem(); + $fs->remove(array_values($this->sessionFiles())); + $fs->remove(array_values($this->apiTokenFiles())); + // Remove the session files' directories if they are now empty. For a session ID beginning with "cli-", + // the directory may also hold another session's SSH certificates. + foreach (glob($this->config->getSessionDir() . '/sess-*', GLOB_ONLYDIR | GLOB_NOSORT) ?: [] as $dir) { + if ((scandir($dir) ?: []) === ['.', '..']) { + $fs->remove($dir); + } + } + } + + private function credentialHelper(): ?Manager + { + $manager = new Manager($this->config); + + return $manager->isSupported() && $manager->isInstalled() ? $manager : null; + } + + /** + * Finds session files, which are named sess-/sess-.json. + * + * @return array Files keyed by session ID. + */ + private function sessionFiles(): array + { + $files = []; + foreach (glob($this->config->getSessionDir() . '/sess-*/sess-*.json', GLOB_NOSORT) ?: [] as $file) { + $id = substr(basename($file, '.json'), strlen('sess-')); + if (basename(dirname($file)) === 'sess-' . $id) { + $files[$id] = $file; + } + } + + return $files; + } + + /** + * Finds API token files, which are named sess-cli-/api-token. + * + * @return array Files keyed by session ID. + */ + private function apiTokenFiles(): array + { + $files = []; + foreach (glob($this->config->getSessionDir() . '/sess-cli-*/api-token', GLOB_NOSORT) ?: [] as $file) { + $files[substr(basename(dirname($file)), strlen('sess-cli-'))] = $file; + } + + return $files; + } +} diff --git a/legacy/src/Command/Auth/LogoutCommand.php b/legacy/src/Command/Auth/LogoutCommand.php deleted file mode 100644 index eb5fe79d3..000000000 --- a/legacy/src/Command/Auth/LogoutCommand.php +++ /dev/null @@ -1,85 +0,0 @@ -addOption('all', 'a', InputOption::VALUE_NONE, 'Log out from all local sessions') - ->addOption('other', null, InputOption::VALUE_NONE, 'Log out from other local sessions'); - } - - protected function execute(InputInterface $input, OutputInterface $output): int - { - // API tokens set via the environment or the config file cannot be - // removed using this command. - // API tokens set via the auth:api-token-login command will be safely - // deleted. - if ($this->api->hasApiToken(false)) { - $this->stdErr->writeln('Warning: an API token is set via config'); - } - - if (Option::bool($input, 'other') && !Option::bool($input, 'all')) { - $currentSessionId = $this->config->getSessionId(); - $this->stdErr->writeln(sprintf('The current session ID is: %s', $currentSessionId)); - $other = \array_filter($this->api->listSessionIds(), fn($sessionId): bool => $sessionId !== $currentSessionId); - if (empty($other)) { - $this->stdErr->writeln('No other sessions exist.'); - return 0; - } - $this->stdErr->writeln(''); - foreach ($other as $sessionId) { - $api = new Api($this->config->withOverrides(['api.session_id' => $sessionId]), null, $output); - $api->logout(); - $this->stdErr->writeln(sprintf('Logged out from session: %s', $sessionId)); - } - $this->stdErr->writeln(''); - $this->stdErr->writeln('All other sessions have been deleted.'); - return 0; - } - - $this->api->logout(); - $this->stdErr->writeln('You are now logged out.'); - $this->sshConfig->deleteSessionConfiguration(); - - // Check for other sessions. - if (Option::bool($input, 'all')) { - $this->api->deleteAllSessions(); - $this->stdErr->writeln(''); - $this->stdErr->writeln('All sessions have been deleted.'); - $this->api->showSessionInfo(true); - return 0; - } - - $this->api->showSessionInfo(true); - - if ($this->api->anySessionsExist()) { - $this->stdErr->writeln(''); - $this->stdErr->writeln(sprintf( - 'Other sessions exist. Log out of all sessions with: %s logout --all', - $this->config->getStr('application.executable'), - )); - } - - return 0; - } -} diff --git a/legacy/src/Command/Auth/PostLoginCommand.php b/legacy/src/Command/Auth/PostLoginCommand.php new file mode 100644 index 000000000..946832789 --- /dev/null +++ b/legacy/src/Command/Auth/PostLoginCommand.php @@ -0,0 +1,33 @@ +login->finalize(); + + return 0; + } +} diff --git a/legacy/src/Exception/GoLoginRequiredException.php b/legacy/src/Exception/GoLoginRequiredException.php new file mode 100644 index 000000000..1bcd3a65b --- /dev/null +++ b/legacy/src/Exception/GoLoginRequiredException.php @@ -0,0 +1,42 @@ + $data + */ + public static function fromData(array $data): self + { + $notice = isset($data['notice']) && is_string($data['notice']) ? $data['notice'] : ''; + $amr = isset($data['amr']) && is_array($data['amr']) ? array_values(array_filter($data['amr'], 'is_string')) : []; + $maxAge = isset($data['max_age']) && is_int($data['max_age']) ? $data['max_age'] : null; + + return new self($notice, $amr, $maxAge, !empty($data['has_api_token'])); + } + + public function toEvent(): LoginRequiredEvent + { + return new LoginRequiredEvent($this->authMethods, $this->maxAge, $this->hasApiToken); + } +} diff --git a/legacy/src/Service/Api.php b/legacy/src/Service/Api.php index 6880cf30b..15c812f4f 100644 --- a/legacy/src/Service/Api.php +++ b/legacy/src/Service/Api.php @@ -9,7 +9,6 @@ use Platformsh\Client\Model\Deployment\Worker; use Platformsh\Client\Model\Project\Capabilities; use GuzzleHttp\HandlerStack; -use Symfony\Component\Filesystem\Filesystem; use GuzzleHttp\Exception\RequestException; use Platformsh\Client\Model\Integration; use Composer\CaBundle\CaBundle; @@ -21,14 +20,10 @@ use GuzzleHttp\Psr7\Uri; use GuzzleHttp\Psr7\UriResolver; use GuzzleHttp\Utils; -use League\OAuth2\Client\Provider\Exception\IdentityProviderException; use League\OAuth2\Client\Token\AccessToken; -use Platformsh\Cli\CredentialHelper\KeyringUnavailableException; -use Platformsh\Cli\CredentialHelper\Manager; -use Platformsh\Cli\CredentialHelper\SessionStorage as CredentialHelperStorage; use Platformsh\Cli\Event\EnvironmentsChangedEvent; use Platformsh\Cli\Event\LoginRequiredEvent; -use Platformsh\Cli\Exception\ProcessFailedException; +use Platformsh\Cli\Exception\GoLoginRequiredException; use Platformsh\Cli\GuzzleDebugMiddleware; use Platformsh\Cli\Model\Route; use Platformsh\Cli\Util\NestedArrayUtil; @@ -54,8 +49,6 @@ use Platformsh\Client\PlatformClient; use Platformsh\Client\Session\Session; use Platformsh\Client\Session\SessionInterface; -use Platformsh\Client\Session\Storage\File; -use Platformsh\Client\Session\Storage\SessionStorageInterface; use Psr\Http\Message\RequestInterface; use Psr\Http\Message\ResponseInterface; use Symfony\Component\Console\Output\ConsoleOutput; @@ -63,7 +56,6 @@ use Symfony\Component\Console\Output\OutputInterface; use Symfony\Component\EventDispatcher\EventDispatcher; use Symfony\Component\EventDispatcher\EventDispatcherInterface; -use Symfony\Component\Process\Exception\ProcessTimedOutException; use Symfony\Contracts\Service\Attribute\Required; /** @@ -84,7 +76,7 @@ class Api private readonly OutputInterface $output; private readonly OutputInterface $stdErr; private readonly TokenConfig $tokenConfig; - private readonly FileLock $fileLock; + private readonly GoAuth $goAuth; private readonly Io $io; /** @@ -118,16 +110,16 @@ class Api private static array $notFound = []; /** - * Session storage, via files or credential helpers. + * The prefix of placeholder refresh tokens, which identify the access token held by the OAuth 2.0 middleware. * - * @see Api::initSessionStorage() + * Tokens are refreshed by the Go wrapper, so these are never sent anywhere. */ - private ?SessionStorageInterface $sessionStorage = null; + private const GO_REFRESH_TOKEN_PREFIX = 'go:'; /** - * Sets whether we are currently verifying login using a test request. + * The last token from the Go wrapper. */ - public bool $inLoginCheck = false; + private static ?AccessToken $goToken = null; /** * Constructor. @@ -137,7 +129,7 @@ class Api * @param OutputInterface|null $output * @param TokenConfig|null $tokenConfig * @param EventDispatcherInterface|null $dispatcher - * @param FileLock|null $fileLock + * @param GoAuth|null $goAuth * @param Io|null $io */ public function __construct( @@ -146,7 +138,7 @@ public function __construct( ?OutputInterface $output = null, ?Io $io = null, ?TokenConfig $tokenConfig = null, - ?FileLock $fileLock = null, + ?GoAuth $goAuth = null, ?EventDispatcherInterface $dispatcher = null, ) { $this->config = $config ?: new Config(); @@ -154,7 +146,7 @@ public function __construct( $this->stdErr = $this->output instanceof ConsoleOutputInterface ? $this->output->getErrorOutput() : $this->output; $this->io = $io ?: new Io($this->output); $this->tokenConfig = $tokenConfig ?: new TokenConfig($this->config); - $this->fileLock = $fileLock ?: new FileLock($this->config); + $this->goAuth = $goAuth ?: new GoAuth($this->config, $this->output); $this->dispatcher = $dispatcher ?: new EventDispatcher(); $this->cache = $cache ?: CacheFactory::createCacheProvider($this->config); } @@ -183,13 +175,17 @@ public function injectListeners( /** * Returns whether the CLI is authenticating using an API token. * - * @param bool $includeStored + * @param bool $includeStored Whether to include an API token saved by auth:api-token-login. * * @return bool */ public function hasApiToken(bool $includeStored = true): bool { - return $this->tokenConfig->getAccessToken() || $this->tokenConfig->getApiToken($includeStored); + if ($this->tokenConfig->getAccessToken() || $this->tokenConfig->getApiToken()) { + return true; + } + + return $includeStored && $this->goAuth->getStatus()['has_stored_api_token']; } /** @@ -201,77 +197,7 @@ public function hasApiToken(bool $includeStored = true): bool */ public function listSessionIds(): array { - $ids = []; - if ($this->sessionStorage instanceof CredentialHelperStorage) { - $ids = $this->sessionStorage->listSessionIds(); - } - $dir = $this->config->getSessionDir(); - $files = glob($dir . '/sess-cli-*', GLOB_NOSORT); - if ($files !== false) { - foreach ($files as $file) { - if (\preg_match('@/sess-cli-([a-z0-9_-]+)@i', $file, $matches)) { - $ids[] = $matches[1]; - } - } - } - $ids = \array_filter($ids, fn($id): bool => !str_starts_with((string) $id, 'api-token-')); - - return \array_unique($ids); - } - - /** - * Checks if any sessions exist (with any session ID). - * - * @return bool - */ - public function anySessionsExist(): bool - { - if ($this->sessionStorage instanceof CredentialHelperStorage && $this->sessionStorage->hasAnySessions()) { - return true; - } - $dir = $this->config->getSessionDir(); - $files = glob($dir . '/sess-cli-*/*', GLOB_NOSORT); - - return !empty($files); - } - - /** - * Logs out of the current session. - */ - public function logout(): void - { - // Delete the stored API token, if any. - $this->tokenConfig->storage()->deleteToken(); - - // Log out in the connector (this clears the "session" and attempts - // to revoke stored tokens). - $this->getClient(false)->getConnector()->logOut(); - - // Clear the cache. - $this->cache->flushAll(); - - // Ensure the session directory is wiped. - $dir = $this->config->getSessionDir(true); - if (is_dir($dir)) { - (new Filesystem())->remove($dir); - } - - // Wipe the client so it is re-initialized when needed. - self::$client = null; - } - - /** - * Deletes all sessions. - */ - public function deleteAllSessions(): void - { - if ($this->sessionStorage instanceof CredentialHelperStorage) { - $this->sessionStorage->deleteAll(); - } - $dir = $this->config->getSessionDir(); - if (is_dir($dir)) { - (new Filesystem())->remove($dir); - } + return $this->goAuth->getStatus()['session_ids']; } /** @@ -296,52 +222,13 @@ private function getConnectorOptions(): array $connectorOptions['user_agent'] = $this->config->getUserAgent(); $connectorOptions['timeout'] = $this->config->getInt('api.default_timeout'); - if ($apiToken = $this->tokenConfig->getApiToken()) { - $connectorOptions['api_token'] = $apiToken; - $connectorOptions['api_token_type'] = 'exchange'; - } elseif ($accessToken = $this->tokenConfig->getAccessToken()) { - $connectorOptions['api_token'] = $accessToken; - $connectorOptions['api_token_type'] = 'access'; - } - $connectorOptions['proxy'] = $this->guzzleProxyConfig(); $connectorOptions['token_url'] = $this->config->get('api.oauth2_token_url'); $connectorOptions['revoke_url'] = $this->config->get('api.oauth2_revoke_url'); - // Acquire a lock to prevent tokens being refreshed at the same time in - // different CLI processes. - $refreshLockName = 'refresh--' . $this->config->getSessionIdSlug(); - $connectorOptions['on_refresh_start'] = function (string $originalRefreshToken) use ($refreshLockName): ?AccessToken { - $this->io->debug('Refreshing access token'); - $this->fileLock->acquireOrWait($refreshLockName, function (): void { - $this->stdErr->writeln('Waiting for token refresh lock', OutputInterface::VERBOSITY_VERBOSE); - }); - // Without the lock, the token could be refreshed or saved over another process's newer token. - if (!$this->fileLock->isHeld($refreshLockName)) { - throw new \RuntimeException('Timed out waiting for another process to refresh the access token. Please try again.'); - } - - // Refresh tokens are single-use, so use the stored token if another process has refreshed it. - $session = $this->getClient(false)->getConnector()->getSession(); - try { - $session->reload(); - } catch (\RuntimeException $e) { - throw $this->convertStorageException($e); - } - $storedToken = $this->tokenFromSession($session); - if ($storedToken && $storedToken->getRefreshToken() !== $originalRefreshToken) { - return $storedToken; - } - - // The middleware refreshes the token, and saves it before calling on_refresh_end. - return null; - }; - $connectorOptions['on_refresh_end'] = function () use ($refreshLockName): void { - $this->fileLock->release($refreshLockName); - }; - - $connectorOptions['on_refresh_error'] = fn(IdentityProviderException $e): ?AccessToken => $this->onRefreshError($e); + // Tokens are refreshed by the Go wrapper. The middleware calls this when its token has expired, or after a 401. + $connectorOptions['on_refresh_start'] = fn(string $refreshToken): AccessToken => $this->refreshFromGo($refreshToken); $connectorOptions['on_step_up_auth_response'] = fn(ResponseInterface $response) => $this->onStepUpAuthResponse($response); @@ -381,17 +268,10 @@ private function caBundlePath(): string return $path; } - private function onStepUpAuthResponse(ResponseInterface $response): ?AccessToken + private function onStepUpAuthResponse(ResponseInterface $response): AccessToken { - if ($this->inLoginCheck) { - return null; - } - $this->io->debug(ApiResponseException::getErrorDetails($response)); - $session = $this->getClient(false)->getConnector()->getSession(); - $previousAccessToken = $session->get('accessToken'); - $body = (array) Utils::jsonDecode((string) $response->getBody(), true); $authMethods = $body['amr'] ?? []; $maxAge = $body['max_age'] ?? null; @@ -399,123 +279,63 @@ private function onStepUpAuthResponse(ResponseInterface $response): ?AccessToken $this->dispatcher->dispatch(new LoginRequiredEvent($authMethods, $maxAge, $this->hasApiToken()), 'login.required'); $this->stdErr->writeln(''); - $session = $this->getClient(false)->getConnector()->getSession(); - $newAccessToken = $this->tokenFromSession($session); - if ($newAccessToken && $newAccessToken->getToken() !== $previousAccessToken) { - return $newAccessToken; - } - return null; + return $this->setGoToken($this->fetchGoToken()); } /** - * Logs out and prompts for re-authentication after a token refresh error. + * Gets a new token from the Go wrapper, for the OAuth 2.0 middleware. * - * @param IdentityProviderException $e - * - * @return AccessToken|null + * @param string $refreshToken The placeholder refresh token of the middleware's current access token. */ - private function onRefreshError(IdentityProviderException $e): ?AccessToken + private function refreshFromGo(string $refreshToken): AccessToken { - if ($this->inLoginCheck) { - return null; + // If the middleware's token has not expired locally, it was rejected by the API. + $current = substr($refreshToken, strlen(self::GO_REFRESH_TOKEN_PREFIX)); + $rejected = null; + if ($current !== '' && self::$goToken?->getToken() === $current && !self::$goToken->hasExpired()) { + $rejected = $current; } - $data = $e->getResponseBody(); - if (!is_array($data) || !isset($data['error'])) { - return null; - } - - $this->io->debug($e->getMessage()); - $this->logout(); - - if ($this->isSsoSessionExpired($data)) { - $this->stdErr->writeln('Your SSO session has expired. You have been logged out.'); - } elseif ($this->isApiTokenInvalid($data)) { - $this->stdErr->writeln('The API token is invalid.'); - } else { - $this->stdErr->writeln('Your session has expired. You have been logged out.'); - } - - $this->stdErr->writeln(''); - - $this->dispatcher->dispatch(new LoginRequiredEvent([], null, $this->hasApiToken()), 'login.required'); - $session = $this->getClient(false)->getConnector()->getSession(); - - return $this->tokenFromSession($session); - } - - /** - * Tests if an HTTP response from refreshing a token indicates that the user's SSO session has expired. - * - * @param array $data - */ - private function isSsoSessionExpired(array $data): bool - { - if (isset($data['error']) && $data['error'] === 'invalid_grant') { - return isset($data['error_description']) - && str_contains((string) $data['error_description'], 'SSO session has expired'); - } - return false; + return $this->setGoToken($this->fetchGoToken($rejected)); } /** - * Tests if an error from refreshing a token indicates that the user's API token is invalid. + * Gets a token from the Go wrapper, offering a login if one is required. * - * @param array $body - */ - private function isApiTokenInvalid(mixed $body): bool - { - if (is_array($body) && isset($body['error']) && $body['error'] === 'invalid_grant') { - return isset($body['error_description']) - && str_contains((string) $body['error_description'], 'API token'); - } - return false; - } - - /** - * Converts a session storage error into a keyring error, if applicable. + * @return array{access_token: string, expires?: int} */ - private function convertStorageException(\RuntimeException $e): \RuntimeException + private function fetchGoToken(?string $rejected = null): array { - if ($this->sessionStorage instanceof CredentialHelperStorage) { - $previous = $e->getPrevious(); - if ($previous instanceof ProcessTimedOutException) { - return KeyringUnavailableException::fromTimeout($previous); - } elseif ($previous instanceof ProcessFailedException) { - return KeyringUnavailableException::fromFailure($previous); + try { + return $this->goAuth->getToken($rejected); + } catch (GoLoginRequiredException $e) { + if ($e->notice !== '') { + $this->stdErr->writeln('' . $e->notice . ''); + $this->stdErr->writeln(''); } + // The listener logs in, or throws a LoginRequiredException. + $this->dispatcher->dispatch($e->toEvent(), 'login.required'); + $this->goAuth->reset(); + + return $this->goAuth->getToken(); } - return $e; } /** - * Loads and returns an AccessToken, if possible, from a session. + * Saves a token from Go, in the form used by the OAuth 2.0 middleware, which refreshes it when Go would. * - * @param SessionInterface $session - * - * @return AccessToken|null + * @param array{access_token: string, expires?: int} $token */ - private function tokenFromSession(SessionInterface $session): ?AccessToken + private function setGoToken(array $token): AccessToken { - if (!$session->get('accessToken')) { - return null; - } - $map = [ - 'accessToken' => 'access_token', - 'expires' => 'expires', - 'refreshToken' => 'refresh_token', - 'scope' => 'scope', - ]; - $tokenData = []; - foreach ($map as $sessionKey => $tokenKey) { - $value = $session->get($sessionKey); - if ($value !== false && $value !== null) { - $tokenData[$tokenKey] = $value; - } - } + $expires = !empty($token['expires']) ? $token['expires'] - 120 : 2147483647; - return new AccessToken($tokenData); + return self::$goToken = new AccessToken([ + 'access_token' => $token['access_token'], + 'expires' => max($expires, time() + 1), + 'refresh_token' => self::GO_REFRESH_TOKEN_PREFIX . $token['access_token'], + ]); } /** @@ -557,30 +377,14 @@ public function getClient(bool $autoLogin = true, bool $reset = false): Platform if (!isset(self::$client) || $reset) { $options = $this->getConnectorOptions(); - $sessionId = $this->config->getSessionId(); - - // Override the session ID if an API token is set. - // This ensures file storage from other credentials will not be - // reused. - if (!empty($options['api_token'])) { - $sessionId = 'api-token-' . \substr(\hash('sha256', (string) $options['api_token']), 0, 32); - } - - // Set up a session to store OAuth2 tokens. - // By default this uses in-memory storage. - $session = new Session($sessionId); - - // Set up persistent session storage - // (unless an access token was set directly). - if (!isset($options['api_token']) || $options['api_token_type'] !== 'access') { - $this->initSessionStorage(); - $this->io->debug('Loading session'); - try { - $session->setStorage($this->sessionStorage); - } catch (\RuntimeException $e) { - throw $this->convertStorageException($e); - } - } + // The session is only kept in memory: tokens are stored by the Go wrapper. It starts with an expired + // placeholder, which the middleware replaces with a token from Go before the first request. + $session = new Session($this->config->getSessionId()); + $this->setSessionToken($session, new AccessToken([ + 'access_token' => 'pending', + 'expires' => time() - 1, + 'refresh_token' => self::GO_REFRESH_TOKEN_PREFIX, + ])); $connector = new Connector($options, $session); @@ -593,13 +397,26 @@ public function getClient(bool $autoLogin = true, bool $reset = false): Platform $this->stdErr->writeln(''); self::$printedApiTokenWarning = true; } + } - if ($autoLogin && !$connector->isLoggedIn()) { - $this->dispatcher->dispatch(new LoginRequiredEvent([], null, $this->hasApiToken()), 'login.required'); - } + $client = self::$client; + + // Get a token now, so that a login is offered if needed, and the token can be read from the session. + $session = $client->getConnector()->getSession(); + if ($autoLogin && $session->get('accessToken') === 'pending') { + $this->setSessionToken($session, self::$goToken !== null && !self::$goToken->hasExpired() + ? self::$goToken + : $this->setGoToken($this->fetchGoToken())); } - return self::$client; + return $client; + } + + private function setSessionToken(SessionInterface $session, AccessToken $token): void + { + $session->set('accessToken', $token->getToken()); + $session->set('expires', $token->getExpires()); + $session->set('refreshToken', $token->getRefreshToken()); } /** @@ -615,31 +432,6 @@ private function onContainer(): bool && getenv($envPrefix . 'TREE_ID') !== false; } - /** - * Initializes session credential storage. - */ - private function initSessionStorage(): void - { - if (!isset($this->sessionStorage)) { - // Attempt to use the docker-credential-helpers. - $manager = new Manager($this->config); - if ($manager->isSupported()) { - if ($manager->isInstalled()) { - $this->io->debug('Using Docker credential helper for session storage'); - } else { - $this->io->debug('Installing Docker credential helper for session storage'); - $manager->install(); - } - $this->sessionStorage = new CredentialHelperStorage($manager, $this->config->getStr('application.slug')); - return; - } - - // Fall back to file storage. - $this->io->debug('Using filesystem for session storage'); - $this->sessionStorage = new File($this->config->getSessionDir()); - } - } - /** * Constructs a stream context for using the API with stream functions. * @@ -1301,7 +1093,7 @@ public static function getNestedProperty(ApiResource $resource, string $property */ public function isLoggedIn(): bool { - return $this->getClient(false)->getConnector()->isLoggedIn(); + return $this->goAuth->getStatus()['logged_in']; } /** @@ -1427,22 +1219,12 @@ public function getAccessToken(bool $forceNew = false): string return $accessToken; } - // Get the access token from the session. - $session = $this->getClient()->getConnector()->getSession(); - $token = $session->get('accessToken'); - $expires = $session->get('expires'); - - // If there is no token, or it has expired, make an API request, which - // automatically obtains a token and saves it to the session. - if (!$token || $expires < time() || $forceNew) { - $this->getUser(null, true); - $newSession = $this->getClient()->getConnector()->getSession(); - if (!$token = $newSession->get('accessToken')) { - throw new \RuntimeException('No access token found'); - } + $this->getClient(); + if ($forceNew || self::$goToken === null || self::$goToken->hasExpired()) { + return $this->setGoToken($this->fetchGoToken($forceNew ? self::$goToken?->getToken() : null))->getToken(); } - return $token; + return self::$goToken->getToken(); } /** diff --git a/legacy/src/Service/AutoLoginListener.php b/legacy/src/Service/AutoLoginListener.php index 1ddb20f28..a7f80def7 100644 --- a/legacy/src/Service/AutoLoginListener.php +++ b/legacy/src/Service/AutoLoginListener.php @@ -16,7 +16,7 @@ public function __construct( private Api $api, - private SubCommandRunner $commandDispatcher, + private GoAuth $goAuth, private Config $config, private InputInterface $input, private QuestionHelper $questionHelper, @@ -52,7 +52,7 @@ public function onLoginRequired(LoginRequiredEvent $event): void } if ($this->questionHelper->confirm('Log in via a browser?')) { $this->stdErr->writeln(''); - $exitCode = $this->commandDispatcher->run('auth:browser-login', $event->getLoginOptions()); + $exitCode = $this->goAuth->login($event->getLoginOptions()); $this->stdErr->writeln(''); $success = $exitCode === 0; } diff --git a/legacy/src/Service/GoAuth.php b/legacy/src/Service/GoAuth.php new file mode 100644 index 000000000..ef351a9d2 --- /dev/null +++ b/legacy/src/Service/GoAuth.php @@ -0,0 +1,183 @@ + + */ + private static array $status = []; + + /** @var array */ + private static array $tokens = []; + + private readonly OutputInterface $stdErr; + + public function __construct(private readonly Config $config, ?OutputInterface $output = null) + { + $output = $output ?: new ConsoleOutput(); + $this->stdErr = $output instanceof ConsoleOutputInterface ? $output->getErrorOutput() : $output; + } + + /** + * Returns an access token and its expiry, refreshing it in Go if needed. + * + * @param string|null $rejected An access token that the API rejected, which Go should replace. + * + * @throws GoLoginRequiredException if login is required + * + * @return array{access_token: string, expires?: int} + */ + public function getToken(?string $rejected = null): array + { + $sessionId = $this->config->getSessionId(); + $cached = self::$tokens[$sessionId] ?? null; + if ($cached !== null && $rejected === null && !$this->expiresSoon($cached)) { + return $cached; + } + // The rejected token is passed via stdin, to keep it out of process listings. + /** @var array{access_token: string, expires?: int} $token */ + $token = $rejected !== null ? $this->run(['token', '--rejected'], $rejected) : $this->run(['token']); + + return self::$tokens[$sessionId] = $token; + } + + /** + * Returns the auth state, without making network requests. + * + * @return array{logged_in: bool, session_ids: string[], has_stored_api_token: bool} + */ + public function getStatus(): array + { + $sessionId = $this->config->getSessionId(); + if (!isset(self::$status[$sessionId])) { + /** @var array{logged_in: bool, session_ids: string[], has_stored_api_token: bool} $status */ + $status = $this->run(['status']); + self::$status[$sessionId] = $status; + } + + return self::$status[$sessionId]; + } + + /** + * Clears cached tokens and state, e.g. after a login. + */ + public function reset(): void + { + self::$status = []; + self::$tokens = []; + } + + /** + * Runs the Go browser login, with the terminal's input and output. + * + * @param array $options Login options, e.g. ['--method' => ['mfa'], '--max-age' => 60]. + * + * @return int The exit code. + */ + public function login(array $options = []): int + { + $command = [$this->executable(), 'auth:browser-login']; + foreach ($options as $option => $value) { + $command[] = $option . '=' . (is_array($value) ? implode(',', $value) : (string) $value); + } + $process = proc_open($command, [STDIN, STDOUT, STDERR], $pipes); + if ($process === false) { + throw new \RuntimeException('Failed to start the login command'); + } + $exitCode = proc_close($process); + $this->reset(); + + return $exitCode; + } + + /** + * @param array{access_token: string, expires?: int} $token + */ + private function expiresSoon(array $token): bool + { + return !empty($token['expires']) && $token['expires'] - 120 < time(); + } + + /** + * @param string[] $args + * @param string|null $input Input for the command's stdin. + * + * @return array + */ + private function run(array $args, ?string $input = null): array + { + $process = new Process(array_merge([$this->executable(), 'auth:internal'], $args), null, $this->env(), $input); + $process->setTimeout(null); + $process->run(); + + $stderr = rtrim($process->getErrorOutput()); + $lines = $stderr === '' ? [] : explode("\n", $stderr); + if ($process->getExitCode() === self::LOGIN_REQUIRED_EXIT_CODE && $lines !== []) { + $data = json_decode((string) array_pop($lines), true); + $this->forwardStderr($lines); + throw GoLoginRequiredException::fromData(is_array($data) ? $data : []); + } + if (!$process->isSuccessful()) { + throw new \RuntimeException('Failed to get authentication details: ' . ($stderr ?: 'exit code ' . $process->getExitCode())); + } + $this->forwardStderr($lines); + $data = json_decode($process->getOutput(), true); + if (!is_array($data)) { + throw new \RuntimeException('Failed to parse authentication details'); + } + + return $data; + } + + /** + * @param string[] $lines + */ + private function forwardStderr(array $lines): void + { + foreach ($lines as $line) { + $this->stdErr->writeln($line, OutputInterface::OUTPUT_RAW); + } + } + + private function executable(): string + { + $var = $this->config->getStr('application.env_prefix') . 'WRAPPER_EXECUTABLE'; + $executable = getenv($var); + if ($executable === false || $executable === '') { + throw new \RuntimeException(sprintf('The %s environment variable is not set. The CLI must be run via its main executable.', $var)); + } + + return $executable; + } + + /** + * Returns environment variables for the Go command, which must use the same session as this process. + * + * @return array + */ + private function env(): array + { + return [$this->config->getStr('application.env_prefix') . 'SESSION_ID' => $this->config->getSessionId()]; + } +} diff --git a/legacy/src/Service/Login.php b/legacy/src/Service/Login.php index 0831f0c9b..bd520222b 100644 --- a/legacy/src/Service/Login.php +++ b/legacy/src/Service/Login.php @@ -13,7 +13,6 @@ private OutputInterface $stdErr; public function __construct( - private Api $api, private Certifier $certifier, private Config $config, private QuestionHelper $questionHelper, @@ -24,14 +23,10 @@ public function __construct( } /** - * Finalizes login: refreshes SSH certificate, prints account information. + * Sets up SSH after a login: host keys, a certificate, and SSH configuration. */ public function finalize(): void { - // Reset the API client so that it will use the new tokens. - $this->api->getClient(false, true); - $this->stdErr->writeln('You are logged in.'); - // Configure SSH host keys. $this->sshConfig->configureHostKeys(); @@ -52,14 +47,6 @@ public function finalize(): void if ($this->sshConfig->configureSessionSsh()) { $this->sshConfig->addUserSshConfig($this->questionHelper); } - - // Show user account info. - $account = $this->api->getMyAccount(true); - $this->stdErr->writeln(sprintf( - "\nUsername: %s\nEmail address: %s", - $account['username'], - $account['email'], - )); } /** diff --git a/legacy/src/Service/TokenConfig.php b/legacy/src/Service/TokenConfig.php index 6b3045dae..1e47c43f6 100644 --- a/legacy/src/Service/TokenConfig.php +++ b/legacy/src/Service/TokenConfig.php @@ -4,34 +4,22 @@ namespace Platformsh\Cli\Service; -use Platformsh\Cli\ApiToken\StorageInterface; -use Platformsh\Cli\ApiToken\Storage; - +/** + * Reads API tokens and access tokens set via config. + * + * Tokens saved by auth:api-token-login are stored by the Go wrapper. + */ readonly class TokenConfig { private Config $config; - private StorageInterface $apiTokenStorage; public function __construct(?Config $config = null) { $this->config = $config ?: new Config(); - $this->apiTokenStorage = Storage::factory($this->config); - } - - public function storage(): StorageInterface - { - return $this->apiTokenStorage; } - public function getApiToken(bool $includeStored = true): ?string + public function getApiToken(): ?string { - if ($includeStored) { - $storedToken = $this->apiTokenStorage->getToken(); - if ($storedToken !== '') { - return $storedToken; - } - } - $token = (string) $this->config->getWithDefault('api.token', ''); if ($token !== '') { return $token; diff --git a/legacy/tests/Command/Auth/BrowserLoginCommandTest.php b/legacy/tests/Command/Auth/BrowserLoginCommandTest.php deleted file mode 100644 index d59627519..000000000 --- a/legacy/tests/Command/Auth/BrowserLoginCommandTest.php +++ /dev/null @@ -1,75 +0,0 @@ -createMock(ClientInterface::class); - $httpClient->expects($this->once()) - ->method('send') - ->willReturnCallback(function (RequestInterface $request, array $options = []) use (&$sentRequest, &$sentOptions): Response { - $sentRequest = $request; - $sentOptions = $options; - return new Response(200, [], '{"access_token": "test-token", "token_type": "bearer"}'); - }); - - $api = $this->createMock(Api::class); - $api->expects($this->once()) - ->method('getExternalHttpClient') - ->willReturn($httpClient); - - $values = [ - 'api.oauth2_token_url' => $tokenUrl, - 'api.oauth2_client_id' => 'test-client-id', - ]; - $config = $this->createMock(Config::class); - $config->method('getStr')->willReturnCallback(fn(string $key): string => $values[$key] ?? ''); - $config->method('get')->willReturnCallback(fn(string $key): string => $values[$key] ?? ''); - - $command = new BrowserLoginCommand( - $api, - $config, - $this->createMock(Io::class), - $this->createMock(Login::class), - $this->createMock(QuestionHelper::class), - $this->createMock(Url::class), - ); - - $method = new \ReflectionMethod($command, 'getAccessToken'); - $token = $method->invoke($command, 'test-code', 'test-verifier', 'http://127.0.0.1:5000'); - - $this->assertSame('test-token', $token['access_token']); - $this->assertInstanceOf(RequestInterface::class, $sentRequest); - $this->assertSame($tokenUrl, (string) $sentRequest->getUri()); - $this->assertSame(['test-client-id', ''], $sentOptions['auth'] ?? null); - } -} diff --git a/legacy/tests/Service/ApiRefreshLockTest.php b/legacy/tests/Service/ApiRefreshLockTest.php deleted file mode 100644 index 497b8a384..000000000 --- a/legacy/tests/Service/ApiRefreshLockTest.php +++ /dev/null @@ -1,169 +0,0 @@ -tempDirSetUp(); - $this->storage = new File($this->config()->getSessionDir()); - } - - public function tearDown(): void - { - // Api caches the client statically: do not leak this session into other tests. - (new \ReflectionProperty(Api::class, 'client'))->setValue(null, null); - } - - public function testUsesTheStoredTokenWhenTheLockIsFree(): void - { - $onRefreshStart = $this->loadStaleSession($this->config()); - - $token = $onRefreshStart('refresh-1'); - - $this->assertInstanceOf(AccessToken::class, $token); - $this->assertSame('access-2', $token->getToken()); - $this->assertSame('refresh-2', $token->getRefreshToken()); - } - - public function testUsesTheStoredTokenAfterWaiting(): void - { - $config = $this->config(); - $onRefreshStart = $this->loadStaleSession($config); - $holder = $this->startLockHolder((string) $this->tempDir, 'refresh--' . $config->getSessionIdSlug(), 1); - - $token = $onRefreshStart('refresh-1'); - \proc_close($holder); - - $this->assertInstanceOf(AccessToken::class, $token); - $this->assertSame('refresh-2', $token->getRefreshToken()); - } - - public function testAllowsARefreshWhenTheStoredTokenIsUnchanged(): void - { - $config = $this->config(); - $this->storage->save('refresh-test', $this->sessionData('access-1', 'refresh-1')); - ['on_refresh_start' => $onRefreshStart, 'on_refresh_end' => $onRefreshEnd] = $this->connector($config)->getConfig(); - $this->assertIsCallable($onRefreshStart); - $this->assertIsCallable($onRefreshEnd); - $lockName = 'refresh--' . $config->getSessionIdSlug(); - $otherProcess = new FileLock($config, 1); - - // The middleware refreshes while the lock is held. - $this->assertNull($onRefreshStart('refresh-1')); - $otherProcess->acquireOrWait($lockName); - $this->assertFalse($otherProcess->isHeld($lockName)); - - $onRefreshEnd('refresh-1'); - $otherProcess->acquireOrWait($lockName); - $this->assertTrue($otherProcess->isHeld($lockName)); - } - - /** - * @return array - */ - public static function timeoutCases(): array - { - return [ - 'unchanged stored token' => [false], - // The lock holder may be about to save an even newer token. - 'newer stored token' => [true], - ]; - } - - #[DataProvider('timeoutCases')] - public function testFailsAfterTimingOut(bool $storedTokenChanged): void - { - $config = $this->config(); - $this->storage->save('refresh-test', $this->sessionData('access-1', 'refresh-1')); - $onRefreshStart = $this->onRefreshStart($config, new FileLock($config, 1)); - if ($storedTokenChanged) { - $this->storage->save('refresh-test', $this->sessionData('access-2', 'refresh-2')); - } - $holder = $this->startLockHolder((string) $this->tempDir, 'refresh--' . $config->getSessionIdSlug()); - - try { - $this->expectExceptionMessage('Timed out waiting for another process to refresh the access token'); - $onRefreshStart('refresh-1'); - } finally { - \proc_terminate($holder, 9); - \proc_close($holder); - } - } - - /** - * @param array $env - */ - private function config(array $env = []): Config - { - return new Config($env + [ - 'PLATFORMSH_CLI_HOME' => (string) $this->tempDir, - 'PLATFORMSH_CLI_SESSION_ID' => 'refresh-test', - 'PLATFORMSH_CLI_API_DISABLE_CREDENTIAL_HELPERS' => '1', - ]); - } - - /** - * Loads the session in memory, then simulates another process refreshing the token. - */ - private function loadStaleSession(Config $config): callable - { - $this->storage->save('refresh-test', $this->sessionData('access-1', 'refresh-1')); - $onRefreshStart = $this->onRefreshStart($config); - $this->storage->save('refresh-test', $this->sessionData('access-2', 'refresh-2')); - - return $onRefreshStart; - } - - private function connector(Config $config, ?FileLock $fileLock = null): Connector - { - $api = new Api($config, new ArrayCache(), new BufferedOutput(), null, null, $fileLock); - $connector = $api->getClient(false, true)->getConnector(); - $this->assertInstanceOf(Connector::class, $connector); - $this->assertSame('refresh-1', $connector->getSession()->get('refreshToken')); - - return $connector; - } - - private function onRefreshStart(Config $config, ?FileLock $fileLock = null): callable - { - $onRefreshStart = $this->connector($config, $fileLock)->getConfig()['on_refresh_start']; - $this->assertIsCallable($onRefreshStart); - - return $onRefreshStart; - } - - /** - * @return array - */ - private function sessionData(string $accessToken, string $refreshToken): array - { - return [ - 'accessToken' => $accessToken, - 'tokenType' => 'bearer', - 'expires' => \time() + 900, - 'refreshToken' => $refreshToken, - ]; - } -} diff --git a/legacy/tests/bootstrap.php b/legacy/tests/bootstrap.php index 46fcedb0d..9e2d78614 100644 --- a/legacy/tests/bootstrap.php +++ b/legacy/tests/bootstrap.php @@ -14,6 +14,9 @@ putenv('PLATFORMSH_CLI_TOKEN='); +// Credentials are managed by the Go wrapper, which is replaced by a stub. +putenv('MOCK_CLI_WRAPPER_EXECUTABLE=' . __DIR__ . '/data/go-auth-stub'); + // Tests run commands in the temporary directory, and the CLI searches parent // directories for a Git repository, so a stray .git would leak into tests. for ($dir = sys_get_temp_dir(); ; $dir = dirname($dir)) { diff --git a/legacy/tests/data/go-auth-stub b/legacy/tests/data/go-auth-stub new file mode 100755 index 000000000..9b603d759 --- /dev/null +++ b/legacy/tests/data/go-auth-stub @@ -0,0 +1,14 @@ +#!/usr/bin/env php + false, 'session_ids' => [], 'has_stored_api_token' => false]) . "\n"; + exit(0); +} +fwrite(STDERR, "{}\n"); +exit(3); diff --git a/pkg/mockapi/api_server.go b/pkg/mockapi/api_server.go index a372783c1..b97ba2726 100644 --- a/pkg/mockapi/api_server.go +++ b/pkg/mockapi/api_server.go @@ -4,6 +4,7 @@ package mockapi import ( "encoding/json" "net/http" + "slices" "strings" "sync" "testing" @@ -21,6 +22,23 @@ type Handler struct { t *testing.T store + + stepUpAMR []string + rejectedTokens []string +} + +// RejectAccessToken makes requests with an access token fail with a 401 error, as if it were revoked. +func (h *Handler) RejectAccessToken(token string) { + h.Lock() + defer h.Unlock() + h.rejectedTokens = append(h.rejectedTokens, token) +} + +// RequireStepUp makes every request fail with a step-up authentication challenge (RFC 9470). +func (h *Handler) RequireStepUp(amr []string) { + h.Lock() + defer h.Unlock() + h.stepUpAMR = amr } func NewHandler(t *testing.T) *Handler { @@ -36,11 +54,27 @@ func NewHandler(t *testing.T) *Handler { authHeader := req.Header.Get("Authorization") require.NotEmpty(t, authHeader) require.True(t, strings.HasPrefix(authHeader, "Bearer ")) + h.RLock() + stepUp := h.stepUpAMR + rejected := slices.Contains(h.rejectedTokens, strings.TrimPrefix(authHeader, "Bearer ")) + h.RUnlock() + if rejected { + w.WriteHeader(http.StatusUnauthorized) + _ = json.NewEncoder(w).Encode(map[string]any{"error": "invalid_token"}) + return + } + if stepUp != nil { + w.Header().Set("WWW-Authenticate", `Bearer error="insufficient_user_authentication"`) + w.WriteHeader(http.StatusUnauthorized) + _ = json.NewEncoder(w).Encode(map[string]any{"amr": stepUp}) + return + } next.ServeHTTP(w, req) }) }) h.Get("/users/me", h.handleUsersMe) + h.Get("/me", h.handleUsersMe) h.Get("/users/{user_id}/extended-access", h.handleUserExtendedAccess) h.Get("/ref/users", h.handleUserRefs) h.Post("/me/verification", func(w http.ResponseWriter, _ *http.Request) { diff --git a/pkg/mockapi/auth_server.go b/pkg/mockapi/auth_server.go index cb0931da6..530ca9fc5 100644 --- a/pkg/mockapi/auth_server.go +++ b/pkg/mockapi/auth_server.go @@ -37,10 +37,63 @@ type AuthServer struct { refreshTokens map[string]*refreshToken refreshDelay time.Duration refreshCount int + refreshRequests int + refreshFailures int + tokenLifetime time.Duration + uniqueTokens bool + accessCount int reuseDetected bool revokedFamilies map[string]bool } +// SetTokenLifetime sets the lifetime of issued access tokens (default: 1 hour). +func (s *AuthServer) SetTokenLifetime(d time.Duration) { + s.refreshMu.Lock() + defer s.refreshMu.Unlock() + s.tokenLifetime = d +} + +// SetUniqueAccessTokens makes the server issue a new access token each time ("access-token-1", "access-token-2", ...), +// instead of always "access-token-1". +func (s *AuthServer) SetUniqueAccessTokens(unique bool) { + s.refreshMu.Lock() + defer s.refreshMu.Unlock() + s.uniqueTokens = unique +} + +func (s *AuthServer) newAccessToken() string { + s.refreshMu.Lock() + defer s.refreshMu.Unlock() + if !s.uniqueTokens { + return accessTokens[0] + } + s.accessCount++ + return fmt.Sprintf("access-token-%d", s.accessCount) +} + +// SetRefreshFailures makes the next n refresh requests fail with a 503 error. +func (s *AuthServer) SetRefreshFailures(n int) { + s.refreshMu.Lock() + defer s.refreshMu.Unlock() + s.refreshFailures = n +} + +// RefreshRequests returns the number of refresh_token grant requests received. +func (s *AuthServer) RefreshRequests() int { + s.refreshMu.Lock() + defer s.refreshMu.Unlock() + return s.refreshRequests +} + +func (s *AuthServer) expiresIn() int { + s.refreshMu.Lock() + defer s.refreshMu.Unlock() + if s.tokenLifetime == 0 { + return 3600 + } + return int(s.tokenLifetime.Seconds()) +} + type refreshToken struct { family string used bool @@ -154,8 +207,8 @@ func NewAuthServer(t *testing.T) *AuthServer { apiToken := req.Form.Get("api_token") if slices.Contains(ValidAPITokens, apiToken) { _ = json.NewEncoder(w).Encode(map[string]any{ - "access_token": accessTokens[0], - "expires_in": 3600, + "access_token": srv.newAccessToken(), + "expires_in": srv.expiresIn(), "token_type": "bearer", "refresh_token": srv.newRefreshToken(), }) @@ -182,25 +235,39 @@ func NewAuthServer(t *testing.T) *AuthServer { return } _ = json.NewEncoder(w).Encode(map[string]any{ - "access_token": accessTokens[0], - "expires_in": 3600, + "access_token": srv.newAccessToken(), + "expires_in": srv.expiresIn(), "token_type": "bearer", "refresh_token": srv.newRefreshToken(), }) case "refresh_token": srv.refreshMu.Lock() + srv.refreshRequests++ delay := srv.refreshDelay + fail := srv.refreshFailures > 0 + if fail { + srv.refreshFailures-- + } srv.refreshMu.Unlock() - time.Sleep(delay) + if fail { + w.WriteHeader(http.StatusServiceUnavailable) + return + } + select { + case <-time.After(delay): + case <-req.Context().Done(): + // The client went away before the token was rotated. + return + } newToken, errCode := srv.rotateRefreshToken(req.Form.Get("refresh_token")) if errCode != "" { writeOAuthError(w, errCode, "The refresh token is invalid.") return } _ = json.NewEncoder(w).Encode(map[string]any{ - "access_token": accessTokens[0], - "expires_in": 3600, + "access_token": srv.newAccessToken(), + "expires_in": srv.expiresIn(), "token_type": "bearer", "refresh_token": newToken, })