diff --git a/README.md b/README.md index aec9cdfe7..d75ddc9d9 100644 --- a/README.md +++ b/README.md @@ -73,6 +73,8 @@ See the [full reference documentation](https://www.digitalocean.com/docs/apis-cl - [Building the Development Version from Source](#building-the-development-version-from-source) - [Dependencies](#dependencies) - [Authenticating with DigitalOcean](#authenticating-with-digitalocean) + - [Signing in with your browser](#signing-in-with-your-browser) + - [Using an API token](#using-an-api-token) - [Logging into multiple DigitalOcean accounts](#logging-into-multiple-digitalocean-accounts) - [Configuring Default Values](#configuring-default-values) - [Environment Variables](#environment-variables) @@ -320,11 +322,37 @@ stable. ## Authenticating with DigitalOcean -To use `doctl`, you need to authenticate with DigitalOcean by providing an access token, which can be created from the [Applications & API](https://cloud.digitalocean.com/account/api/tokens) section of the Control Panel. You can learn how to generate a token by following the [DigitalOcean API guide](https://www.digitalocean.com/community/tutorials/how-to-use-the-digitalocean-api-v2). +`doctl` can authenticate in two ways: sign in through your browser with `doctl auth login`, or paste an API token with `doctl auth init`. Either way, the credentials are saved to an authentication context in the `doctl` configuration file. Docker users will have to use the `DIGITALOCEAN_ACCESS_TOKEN` environmental variable to authenticate, as explained in the Installation section of this document. -If you're not using Docker to run `doctl`, authenticate with the `auth init` command. +### Signing in with your browser + +`doctl auth login` signs you in without creating or pasting a token. It uses the OAuth 2.1 authorization code flow with PKCE, signing in as the `doctl` application DigitalOcean publishes, so there is nothing to register first. + +``` +doctl auth login +``` + +`doctl` opens your browser and waits while you approve the request. The authorization link is valid for 5 minutes; change that with `--timeout`. If the browser does not open, `doctl` prints the link so you can visit it yourself. On a machine with no browser, use `--no-browser` to print the link instead of opening one. + +By default you choose the permissions to grant on the authorization screen. To request them up front, pass `--scopes`: + +``` +doctl auth login --scopes "read write" +``` + +`--scopes "read write"` requests the same access a full-permission API token has. You can also name individual permissions, such as `--scopes "droplet:read account:read"`. Add `--save-scope` to reuse those scopes on later sign-ins, and `--scopes "" --save-scope` to clear the saved default. + +Access tokens issued this way are short-lived. `doctl` saves a refresh token alongside the access token and renews it automatically when you run a command, so you do not have to sign in again when the access token expires. Sign in again to change the granted scopes, or if the session is revoked. + +Without `--context`, signing in replaces the credentials in the `default` context and selects it. Pass `--context ` to sign in to a named context instead, leaving the context you were using unchanged. + +### Using an API token + +To authenticate with an API token instead, create one from the [Applications & API](https://cloud.digitalocean.com/account/api/tokens) section of the Control Panel. You can learn how to generate a token by following the [DigitalOcean API guide](https://www.digitalocean.com/community/tutorials/how-to-use-the-digitalocean-api-v2). + +Authenticate with the `auth init` command. ``` doctl auth init @@ -348,7 +376,7 @@ This will create the necessary directory structure and configuration file to sto `doctl` allows you to log in to multiple DigitalOcean accounts at the same time and easily switch between them with the use of authentication contexts. -By default, a context named `default` is used. To create a new context, run `doctl auth init --context `. You may also pass the new context's name using the `DIGITALOCEAN_CONTEXT` [environment variable](#environment-variables). You will be prompted for your API access token which will be associated with the new context. +By default, a context named `default` is used. To create a new context, run `doctl auth login --context ` to sign in through your browser, or `doctl auth init --context ` to be prompted for an API token. You may also pass the new context's name using the `DIGITALOCEAN_CONTEXT` [environment variable](#environment-variables). The credentials are associated with the new context, and `doctl auth login` also selects it. To use a non-default context, pass the context name to any `doctl` command. For example: @@ -356,7 +384,7 @@ To use a non-default context, pass the context name to any `doctl` command. For doctl compute droplet list --context ``` -To set a new default context, run `doctl auth switch --context `. This command will save the current context to the config file and use it for all commands by default if a context is not specified. +To set a new default context, run `doctl auth switch --context `. This command will save the current context to the config file and use it for all commands by default if a context is not specified. To see the contexts you have, run `doctl auth list`; the selected one is marked `(current)`. The `--access-token` flag or `DIGITALOCEAN_ACCESS_TOKEN` [environment variable](#environment-variables) are acknowledged only if the `default` context is used. Otherwise, they will have no effect on what API access token is used. To temporarily override the access token if a different context is set as default, use `doctl --context default --access-token your_DO_token ...`. diff --git a/args.go b/args.go index 4f734e0d2..55b4b6151 100644 --- a/args.go +++ b/args.go @@ -747,6 +747,21 @@ const ( // ArgTokenValidationServer is the server used to validate an OAuth token ArgTokenValidationServer = "token-validation-server" + // ArgOAuthServer is the OAuth authorization server used to sign in. + ArgOAuthServer = "oauth-server" + // ArgOAuthClientID overrides the OAuth application doctl signs in as. + ArgOAuthClientID = "client-id" + // ArgOAuthScopes is the space-separated list of scopes requested when signing in. + ArgOAuthScopes = "scopes" + // ArgOAuthSaveScope persists --scopes as the default for later logins. + ArgOAuthSaveScope = "save-scope" + // ArgOAuthCallbackPort is the local port that receives the OAuth redirect. + ArgOAuthCallbackPort = "callback-port" + // ArgOAuthNoBrowser prints the authorization URL instead of opening a browser. + ArgOAuthNoBrowser = "no-browser" + // ArgOAuthTimeout bounds how long to wait for browser authorization, in seconds. + ArgOAuthTimeout = "timeout" + // ArgGPUs specifies to list GPU Droplets ArgGPUs = "gpus" diff --git a/commands/auth.go b/commands/auth.go index cc91caf26..5b901f4a3 100644 --- a/commands/auth.go +++ b/commands/auth.go @@ -26,6 +26,7 @@ import ( "github.com/digitalocean/doctl" "github.com/digitalocean/doctl/commands/charm/input" "github.com/digitalocean/doctl/commands/charm/template" + "github.com/digitalocean/doctl/internal/oauth" "github.com/spf13/cobra" "github.com/spf13/viper" @@ -94,22 +95,24 @@ func Auth() *Command { Command: &cobra.Command{ Use: "auth", Short: "Display commands for authenticating doctl with an account", - Long: `The ` + "`" + `doctl auth` + "`" + ` commands allow you to authenticate doctl for use with your DigitalOcean account using tokens that you generate in the control panel at https://cloud.digitalocean.com/account/api/tokens. + Long: `The ` + "`" + `doctl auth` + "`" + ` commands allow you to authenticate doctl for use with your DigitalOcean account, either by signing in through your browser with OAuth or with tokens that you generate in the control panel at https://cloud.digitalocean.com/account/api/tokens. -If you work with a just one account, call ` + "`" + `doctl auth init` + "`" + ` and supply the token when prompted. This creates an authentication context named ` + "`" + `default` + "`" + `. +The quickest way to get started is ` + "`" + `doctl auth login` + "`" + `, which uses the OAuth 2.1 authorization code flow with PKCE. It opens your browser so you can approve access, then saves the access token to the ` + "`" + `default` + "`" + ` authentication context. Tokens issued this way are short-lived; doctl stores a refresh token and renews the access token automatically, so you do not have to return to the control panel when it expires. Use ` + "`" + `--scopes` + "`" + ` to request permissions up front, such as ` + "`" + `--scopes "read write"` + "`" + `, or omit it to choose them on the authorization screen. -To switch between multiple DigitalOcean accounts, including team accounts, create named contexts using ` + "`" + `doctl auth init --context ` + "`" + `, then providing the applicable token when prompted. This saves the token under the name you provide. To switch between contexts, use ` + "`" + `doctl auth switch --context ` + "`" + `. +If you prefer a personal access token instead, call ` + "`" + `doctl auth init` + "`" + ` and supply the token when prompted. This creates an authentication context named ` + "`" + `default` + "`" + `. -To remove accounts from the configuration file, run ` + "`" + `doctl auth remove --context ` + "`" + `. This removes the token under the name you provide.`, +To use multiple DigitalOcean accounts, including team accounts, create a named context with ` + "`" + `doctl auth login --context ` + "`" + ` or ` + "`" + `doctl auth init --context ` + "`" + `. The credentials are saved under the name you provide. To change which context later commands use, run ` + "`" + `doctl auth switch --context ` + "`" + `. + +To remove an account from the configuration file, run ` + "`" + `doctl auth remove --context ` + "`" + `. This deletes the credentials saved under the name you provide.`, GroupID: configureDoctlGroup, }, } - cmdAuthInit := cmdBuilderWithInit(cmd, RunAuthInit(retrieveUserTokenFromCommandLine), "init", "Initialize doctl to use a specific account", `This command allows you to initialize doctl with a token that allows it to query and manage your account details and resources. + cmdAuthInit := cmdBuilderWithInit(cmd, RunAuthInit(retrieveUserTokenFromCommandLine), "init", "Initialize doctl to use a specific account", `This command initializes doctl with an API token, allowing it to query and manage your account details and resources. -The command requires and API token to authenticate, which you can generate in the control panel at https://cloud.digitalocean.com/account/api/tokens. +The command requires an API token to authenticate, which you can generate in the control panel at https://cloud.digitalocean.com/account/api/tokens. To sign in through your browser instead, see the help for `+"`"+`doctl auth login`+"`"+`. -The `+"`"+`--context`+"`"+` flag allows you to add authentication for multiple accounts and then switch between them as needed. Provide a case-sensitive name for the context and then enter the API token you want use for that context when prompted. You can switch authentication contexts using `+"`"+`doctl auth switch`+"`"+`, which re-initializes doctl. You can also provide the `+"`"+`--context`+"`"+` flag when using any doctl command to specify the auth context for that command. This enables you to use multiple DigitalOcean accounts with doctl, or tokens that have different authentication scopes. +The `+"`"+`--context`+"`"+` flag allows you to add authentication for multiple accounts and then switch between them as needed. Provide a case-sensitive name for the context and then enter the API token you want to use for that context when prompted. You can switch authentication contexts using `+"`"+`doctl auth switch`+"`"+`, which re-initializes doctl. You can also provide the `+"`"+`--context`+"`"+` flag when using any doctl command to specify the auth context for that command. This enables you to use multiple DigitalOcean accounts with doctl, or tokens that have different authentication scopes. If the `+"`"+`--context`+"`"+` flag is not specified, doctl creates a default authentication context named `+"`"+`default`+"`"+`. @@ -117,38 +120,54 @@ You can use doctl without initializing it by adding the `+"`"+`--access-token`+" AddStringFlag(cmdAuthInit, doctl.ArgTokenValidationServer, "", TokenValidationServer, "The server used to validate a token") cmdAuthInit.Example = `The following example initializes doctl with a token for a single account with the context ` + "`" + `your-team` + "`" + `: doctl auth init --context your-team` - cmdAuthSwitch := cmdBuilderWithInit(cmd, RunAuthSwitch, "switch", "Switch between authentication contexts", `This command allows you to switch between authentication contexts you've already created. + cmdAuthLogin := cmdBuilderWithInit(cmd, RunAuthLogin, "login", "Sign in to DigitalOcean via OAuth 2.1 in your browser", `This command signs doctl in to your DigitalOcean account in the browser. You approve access there instead of creating and pasting an API token. + +doctl signs in as the doctl application DigitalOcean publishes, so there is nothing to register first. It opens your browser, and after you approve the request it saves the access token and uses that authentication context for later commands. + +The access token is short-lived. doctl stores a refresh token alongside it and renews the access token the next time you run a command. + +With no `+"`"+`--scopes`+"`"+`, you choose the permissions on the authorization screen. `+"`"+`--scopes "read write"`+"`"+` requests the same access as a full-permission API token. You can also name individual permissions, such as `+"`"+`--scopes "droplet:read account:read"`+"`"+`. Add `+"`"+`--save-scope`+"`"+` to reuse those scopes on later logins. `+"`"+`--scopes "" --save-scope`+"`"+` clears the saved default. + +`+"`"+`--context `+"`"+` signs in to that context, creating it when needed, and later commands use it. This keeps several accounts or teams side by side. Without `+"`"+`--context`+"`"+`, the sign-in replaces the credentials in the `+"`"+`default`+"`"+` context and switches to it. The context you were using is left unchanged. + +To sign in with an API token instead, see the help for `+"`"+`doctl auth init`+"`"+`.`, Writer, false) + AddStringFlag(cmdAuthLogin, doctl.ArgOAuthServer, "", oauth.DefaultIssuer, "The OAuth authorization server to sign in to") + AddStringFlag(cmdAuthLogin, doctl.ArgOAuthClientID, "", oauth.DefaultClientID, "The OAuth application to sign in as. Defaults to the doctl application and only needs changing for a non-production authorization server") + AddStringFlag(cmdAuthLogin, doctl.ArgOAuthScopes, "", defaultOAuthScopes, "A space-separated list of scopes to request, such as \"read write\" or \"droplet:read account:read\". When omitted, uses any default saved with --save-scope; otherwise you choose the permissions to grant in your browser") + AddBoolFlag(cmdAuthLogin, doctl.ArgOAuthSaveScope, "", false, "Save --scopes as the default for later logins. Use with --scopes \"\" to clear the saved default") + AddIntFlag(cmdAuthLogin, doctl.ArgOAuthCallbackPort, "", 0, "The local port to listen on for the authorization redirect. Defaults to an unused port") + AddBoolFlag(cmdAuthLogin, doctl.ArgOAuthNoBrowser, "", false, "Print the authorization URL instead of opening a browser") + AddDurationFlag(cmdAuthLogin, doctl.ArgOAuthTimeout, "", oauth.DefaultLoginTimeout, "How long the authorization link stays valid while doctl waits for you to approve the request in your browser") + cmdAuthLogin.Example = `The following example signs in to the context ` + "`" + `your-team` + "`" + `, requests full read and write access, and saves those scopes for later logins: doctl auth login --context your-team --scopes "read write" --save-scope` + + cmdAuthSwitch := cmdBuilderWithInit(cmd, RunAuthSwitch, "switch", "Switch between authentication contexts", `This command changes which authentication context later commands use. The context must already exist. To see a list of available authentication contexts, call `+"`"+`doctl auth list`+"`"+`. -For details on creating an authentication context, see the help for `+"`"+`doctl auth init`+"`"+`.`, Writer, false) +To create a context, see the help for `+"`"+`doctl auth login`+"`"+` or `+"`"+`doctl auth init`+"`"+`.`, Writer, false) cmdAuthSwitch.AddValidArgsFunc(authContextListValidArgsFunc) cmdAuthSwitch.Example = `The following example switches to the context ` + "`" + `your-team` + "`" + `: doctl auth switch --context your-team` - cmdAuthRemove := cmdBuilderWithInit(cmd, RunAuthRemove, "remove --context ", "Remove authentication contexts ", `This command allows you to remove authentication contexts you've already created. + cmdAuthRemove := cmdBuilderWithInit(cmd, RunAuthRemove, "remove --context ", "Remove authentication contexts", `This command removes an authentication context, deleting the credentials saved under that name. Removing a context does not change which context is selected. To see a list of available authentication contexts, call `+"`"+`doctl auth list`+"`"+`. -For details on creating an authentication context, see the help for `+"`"+`doctl auth init`+"`"+`.`, Writer, false) +To create a context, see the help for `+"`"+`doctl auth login`+"`"+` or `+"`"+`doctl auth init`+"`"+`.`, Writer, false) cmdAuthRemove.AddValidArgsFunc(authContextListValidArgsFunc) cmdAuthRemove.Example = `The following example removes the context ` + "`" + `your-team` + "`" + `: doctl auth remove --context your-team` - cmdAuthList := cmdBuilderWithInit(cmd, RunAuthList, "list", "List available authentication contexts", `List named authentication contexts that you created with `+"`"+`doctl auth init`+"`"+`. - -To switch between the contexts use `+"`"+`doctl auth switch --context `+"`"+`, where `+"`"+``+"`"+` is one of the contexts listed. + cmdAuthList := cmdBuilderWithInit(cmd, RunAuthList, "list", "List available authentication contexts", `This command lists the authentication contexts you created with `+"`"+`doctl auth login`+"`"+` or `+"`"+`doctl auth init`+"`"+`. The context later commands use is marked `+"`"+`(current)`+"`"+`. -To create new contexts, see the help for `+"`"+`doctl auth init`+"`"+`.`, Writer, false, aliasOpt("ls")) +To change which context is used, run `+"`"+`doctl auth switch --context `+"`"+`, where `+"`"+``+"`"+` is one of the contexts listed.`, Writer, false, aliasOpt("ls")) // The command runner expects that any command named "list" accepts a // format flag, so we include here despite only supporting text output for // this command. AddStringFlag(cmdAuthList, doctl.ArgFormat, "", "", "Columns for output in a comma-separated list. Possible values: `text`") - cmdAuthList.Example = `The following example lists the available contexts with the ` + "`" + `--format` + "`" + ` flag: doctl auth list` + cmdAuthList.Example = `The following example lists the available authentication contexts: doctl auth list` - cmdAuthToken := cmdBuilderWithInit(cmd, RunAuthToken, "token", "Display current authentication context API token", `Display the current authentication context's token that you created with `+"`"+`doctl auth init`+"`"+`. + cmdAuthToken := cmdBuilderWithInit(cmd, RunAuthToken, "token", "Display current authentication context API token", `This command prints the access token for the current authentication context, whether it came from `+"`"+`doctl auth login`+"`"+` or `+"`"+`doctl auth init`+"`"+`. Treat it like a password: anyone holding it can act on your account. -To switch between the contexts use `+"`"+`doctl auth switch --context `+"`"+`, where `+"`"+``+"`"+` is one of the contexts from: `+"`"+`doctl auth list`+"`"+` - -To create new contexts, see the help for `+"`"+`doctl auth init`+"`"+`.`, Writer, false, aliasOpt("t")) +To print the token for a different context, add `+"`"+`--context `+"`"+`, where `+"`"+``+"`"+` is one of the contexts from `+"`"+`doctl auth list`+"`"+`.`, Writer, false, aliasOpt("t")) cmdAuthToken.Example = `The following example displays the token of the current context: doctl auth token` return cmd @@ -175,6 +194,10 @@ func RunAuthInit(retrieveUserTokenFunc func() (string, error)) func(c *CmdConfig } c.setContextAccessToken(token) + // An explicitly supplied token replaces whatever credential the + // context held, including a browser sign-in that would otherwise be + // refreshed over the top of it. + removeOAuthTokenState(context) template.Render(c.Out, `{{nl}}Validating token... `, nil) @@ -213,6 +236,8 @@ func RunAuthRemove(c *CmdConfig) error { return fmt.Errorf("Context not found") } + removeOAuthTokenState(context) + fmt.Println("Context deleted successfully") return writeConfig() diff --git a/commands/auth_login.go b/commands/auth_login.go new file mode 100644 index 000000000..e88ca3bf2 --- /dev/null +++ b/commands/auth_login.go @@ -0,0 +1,516 @@ +/* +Copyright 2018 The Doctl Authors All rights reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package commands + +import ( + "context" + "errors" + "fmt" + "net/http" + "os" + "strings" + "time" + + "github.com/spf13/viper" + + "github.com/digitalocean/doctl" + "github.com/digitalocean/doctl/commands/charm/template" + "github.com/digitalocean/doctl/internal/oauth" +) + +const ( + // oauthTokensConfigKey holds the per-context OAuth session state used to + // refresh access tokens. + oauthTokensConfigKey = "oauth-tokens" + // oauthDefaultScopesConfigKey holds the optional default scopes reused on + // later logins when the user passed --save-scope. + oauthDefaultScopesConfigKey = "oauth-default-scopes" + + // defaultOAuthScopes is empty on purpose: when doctl sends no scope, the + // authorization screen decides what the token is granted, which lets the + // user pick the permissions they want rather than the CLI assuming them. + defaultOAuthScopes = "" + + // oauthRefreshSkew refreshes an access token slightly before it expires so + // a command that takes a moment to reach the API is not rejected. + oauthRefreshSkew = 2 * time.Minute + + // oauthHTTPTimeout bounds each individual call to the authorization + // server. It does not bound the wait for the user's browser. + oauthHTTPTimeout = 30 * time.Second +) + +// oauthTokenState is the per-context OAuth session persisted in the config +// file. The access token itself lives with the rest of the context's +// credentials; only what is needed to renew it is kept here. +type oauthTokenState struct { + Issuer string + ClientID string + TokenEndpoint string + RefreshToken string + Scope string + ExpiresAt time.Time +} + +func (s *oauthTokenState) needsRefresh(now time.Time) bool { + if s == nil || s.RefreshToken == "" { + return false + } + // An unknown expiry means the token was issued by a server that does not + // report one; leave it alone and let the API reject it if it is stale. + if s.ExpiresAt.IsZero() { + return false + } + + return !now.Add(oauthRefreshSkew).Before(s.ExpiresAt) +} + +// RunAuthLogin authenticates doctl using the OAuth 2.1 authorization code flow +// with PKCE. Without --context, the access token replaces the one in the +// default context and later commands use that context. +func RunAuthLogin(c *CmdConfig) error { + namedContext := strings.TrimSpace(Context) != "" + authContext := doctl.ArgDefaultContext + if namedContext { + authContext = currentAuthContext() + } else { + // setContextAccessToken follows the global --context flag, then the + // selected context. Pin the flag to default so this login does not + // replace the context the user is working in. + Context = doctl.ArgDefaultContext + } + + issuer, err := oauthServer(c) + if err != nil { + return err + } + clientID, err := oauthClientID(c) + if err != nil { + return err + } + scopes, err := oauthScopes(c) + if err != nil { + return err + } + saveScope, err := oauthSaveScope(c) + if err != nil { + return err + } + if saveScope && !c.Doit.IsSet(doctl.ArgOAuthScopes) { + return errors.New("--save-scope requires --scopes") + } + port, err := oauthCallbackPort(c) + if err != nil { + return err + } + noBrowser, err := oauthNoBrowser(c) + if err != nil { + return err + } + timeout, err := oauthTimeout(c) + if err != nil { + return err + } + // writeConfig persists every viper setting, so login-only flags from this + // invocation must not stick around for later runs. + clearEphemeralOAuthLoginFlags(c) + + ctx := context.Background() + httpClient := &http.Client{Timeout: oauthHTTPTimeout} + metadata := oauth.ServerMetadataFor(issuer) + + // --context names the context to update. Without it, login replaces the + // credentials in the default context and switches to it. + if !namedContext { + template.Render(c.Out, `{{nl}}Signing in to the {{highlight .}} context. Credentials saved there will be replaced.{{nl}}`, authContext) + template.Render(c.Out, `{{muted "Later commands will use it. To sign in to another context, cancel and run"}} {{highlight "doctl auth login --context "}}{{nl}}{{nl}}`, nil) + } + + token, err := oauth.Login(ctx, oauth.LoginOptions{ + Metadata: metadata, + ClientID: clientID, + Scopes: strings.Fields(scopes), + Port: port, + Timeout: timeout, + HTTPClient: httpClient, + NoBrowser: noBrowser, + OnAuthorizationURL: authorizationURLPrinter(c, noBrowser, timeout), + }) + if errors.Is(err, oauth.ErrInvalidClient) { + template.Render(c.Out, `{{error crossmark}}{{nl}}{{nl}}`, nil) + return fmt.Errorf("%s rejected the doctl application (client ID %s): %s", metadata.Issuer, clientID, err) + } + if err != nil { + template.Render(c.Out, `{{error crossmark}}{{nl}}{{nl}}`, nil) + return fmt.Errorf("Unable to authenticate with DigitalOcean: %s", err) + } + + template.Render(c.Out, `{{success checkmark}}{{nl}}{{nl}}`, nil) + + c.setContextAccessToken(token.AccessToken) + if token.RefreshToken == "" { + // Without a refresh token there is nothing to renew later, so make + // sure an earlier session is not left behind to be refreshed over the + // top of this one. + removeOAuthTokenState(authContext) + } else { + storeOAuthTokenState(authContext, &oauthTokenState{ + Issuer: metadata.Issuer, + ClientID: clientID, + TokenEndpoint: metadata.TokenEndpoint, + RefreshToken: token.RefreshToken, + Scope: token.Scope, + ExpiresAt: token.Expiry, + }) + } + if saveScope { + storeOAuthDefaultScopes(scopes) + } + + // The context just signed in to is the one later commands should use, + // including a named context passed with --context. + viper.Set(doctl.ArgContext, authContext) + + if err := writeConfig(); err != nil { + return err + } + + displayOAuthLoginSummary(c, authContext, token) + if saveScope { + displayOAuthDefaultScopesSaved(c, scopes) + } + + return nil +} + +// oauthServer returns the authorization server for this invocation. +// A value left in the config file from an earlier login is ignored. +func oauthServer(c *CmdConfig) (string, error) { + if !c.Doit.IsSet(doctl.ArgOAuthServer) { + return oauth.DefaultIssuer, nil + } + return c.Doit.GetString(c.NS, doctl.ArgOAuthServer) +} + +// oauthClientID returns the OAuth application doctl signs in as. It is the +// application DigitalOcean registered for doctl unless --client-id names +// another one, which is what a non-production authorization server needs. +func oauthClientID(c *CmdConfig) (string, error) { + if !c.Doit.IsSet(doctl.ArgOAuthClientID) { + return oauth.DefaultClientID, nil + } + + clientID, err := c.Doit.GetString(c.NS, doctl.ArgOAuthClientID) + if err != nil { + return "", err + } + if strings.TrimSpace(clientID) == "" { + return "", errors.New("--client-id cannot be empty") + } + + return clientID, nil +} + +// oauthScopes returns the scopes requested for this invocation. An explicit +// --scopes wins; otherwise any default saved with --save-scope is used. A bare +// --scopes value left in the config file from an earlier login is ignored. +func oauthScopes(c *CmdConfig) (string, error) { + if c.Doit.IsSet(doctl.ArgOAuthScopes) { + return c.Doit.GetString(c.NS, doctl.ArgOAuthScopes) + } + if saved := loadOAuthDefaultScopes(); saved != "" { + return saved, nil + } + return defaultOAuthScopes, nil +} + +// oauthSaveScope reports whether this invocation passed --save-scope. +// A value left in the config file from an earlier login is ignored. +func oauthSaveScope(c *CmdConfig) (bool, error) { + if !c.Doit.IsSet(doctl.ArgOAuthSaveScope) { + return false, nil + } + return c.Doit.GetBool(c.NS, doctl.ArgOAuthSaveScope) +} + +// oauthCallbackPort returns the local callback port for this invocation. +// A value left in the config file from an earlier login is ignored. +func oauthCallbackPort(c *CmdConfig) (int, error) { + if !c.Doit.IsSet(doctl.ArgOAuthCallbackPort) { + return 0, nil + } + return c.Doit.GetInt(c.NS, doctl.ArgOAuthCallbackPort) +} + +// oauthNoBrowser reports whether this invocation passed --no-browser. +// A value left in the config file from an earlier login is ignored. +func oauthNoBrowser(c *CmdConfig) (bool, error) { + if !c.Doit.IsSet(doctl.ArgOAuthNoBrowser) { + return false, nil + } + return c.Doit.GetBool(c.NS, doctl.ArgOAuthNoBrowser) +} + +// oauthTimeout returns how long to wait for browser authorization. +// A value left in the config file from an earlier login is ignored. +func oauthTimeout(c *CmdConfig) (time.Duration, error) { + if !c.Doit.IsSet(doctl.ArgOAuthTimeout) { + return oauth.DefaultLoginTimeout, nil + } + return c.Doit.GetDuration(c.NS, doctl.ArgOAuthTimeout) +} + +// clearEphemeralOAuthLoginFlags resets login-only flags so writeConfig does +// not persist them into the doctl configuration file. +func clearEphemeralOAuthLoginFlags(c *CmdConfig) { + viper.Set(c.NS+"."+doctl.ArgOAuthServer, oauth.DefaultIssuer) + viper.Set(c.NS+"."+doctl.ArgOAuthClientID, oauth.DefaultClientID) + viper.Set(c.NS+"."+doctl.ArgOAuthScopes, defaultOAuthScopes) + viper.Set(c.NS+"."+doctl.ArgOAuthSaveScope, false) + viper.Set(c.NS+"."+doctl.ArgOAuthCallbackPort, 0) + viper.Set(c.NS+"."+doctl.ArgOAuthNoBrowser, false) + viper.Set(c.NS+"."+doctl.ArgOAuthTimeout, oauth.DefaultLoginTimeout) +} + +// loadOAuthDefaultScopes returns the scopes saved with --save-scope, if any. +func loadOAuthDefaultScopes() string { + return viper.GetString(oauthDefaultScopesConfigKey) +} + +// storeOAuthDefaultScopes persists the default scopes for later logins. An +// empty value clears any previously saved default. +func storeOAuthDefaultScopes(scopes string) { + if scopes == "" { + viper.Set(oauthDefaultScopesConfigKey, nil) + return + } + viper.Set(oauthDefaultScopesConfigKey, scopes) +} + +func authorizationURLPrinter(c *CmdConfig, noBrowser bool, timeout time.Duration) func(string) { + return func(authURL string) { + prompt := authorizationPrompt{URL: authURL, Timeout: formatLoginTimeout(timeout)} + if noBrowser { + template.Render(c.Out, `{{nl}}Visit the link below to authorize doctl.{{nl}}{{nl}}The link below expires after {{highlight .Timeout}}.{{nl}}{{nl}} {{underline .URL}}{{nl}}{{nl}}Waiting for authorization... `, prompt) + return + } + + template.Render(c.Out, + `{{nl}}Opening your browser to authorize doctl. If it does not open, visit the link below.{{nl}}{{nl}}The link below expires after {{highlight .Timeout}}.{{nl}}{{nl}} {{underline .URL}}{{nl}}{{nl}}Waiting for authorization... `, + prompt, + ) + } +} + +type authorizationPrompt struct { + URL string + Timeout string +} + +// formatLoginTimeout renders a wait as words, such as "5 minutes". +func formatLoginTimeout(d time.Duration) string { + d = d.Truncate(time.Second) + if d > 0 && d%time.Minute == 0 { + minutes := int(d / time.Minute) + if minutes == 1 { + return "1 minute" + } + return fmt.Sprintf("%d minutes", minutes) + } + + seconds := int(d / time.Second) + if seconds == 1 { + return "1 second" + } + return fmt.Sprintf("%d seconds", seconds) +} + +func displayOAuthLoginSummary(c *CmdConfig, authContext string, token *oauth.Token) { + who := token.Info.Email + if who == "" { + who = token.Info.Name + } + if who != "" { + template.Render(c.Out, `Signed in as {{highlight .}}{{nl}}`, who) + } + if token.Info.TeamName != "" { + template.Render(c.Out, `Team: {{highlight .}}{{nl}}`, token.Info.TeamName) + } + + template.Render(c.Out, `Saved to the {{highlight .}} context. Later commands will use it.{{nl}}`, authContext) + + if token.RefreshToken != "" { + template.Render(c.Out, `{{nl}}{{muted "This token renews automatically."}}{{nl}}`, nil) + template.Render(c.Out, `{{muted "Run"}} {{highlight "doctl auth login"}} {{muted "again to change its scopes, or if the session is revoked."}}{{nl}}`, nil) + } +} + +func displayOAuthDefaultScopesSaved(c *CmdConfig, scopes string) { + if scopes == "" { + template.Render(c.Out, `{{muted "Cleared the saved default scopes."}}{{nl}}`, nil) + return + } + template.Render(c.Out, `{{muted "Saved default scopes for later logins:"}} {{highlight .}}{{nl}}`, scopes) +} + +// refreshExpiredOAuthToken renews the current context's access token when it +// was obtained with doctl auth login and has expired. It is a no-op for +// contexts authenticated with a personal access token. +func refreshExpiredOAuthToken(c *CmdConfig) error { + if accessTokenOverridden() { + return nil + } + + authContext := currentAuthContext() + state := loadOAuthTokenState(authContext) + if !state.needsRefresh(time.Now()) { + return nil + } + + endpoint := state.TokenEndpoint + if endpoint == "" { + endpoint = oauth.ServerMetadataFor(state.Issuer).TokenEndpoint + } + + httpClient := &http.Client{Timeout: oauthHTTPTimeout} + token, err := oauth.RefreshToken(context.Background(), httpClient, endpoint, state.ClientID, state.RefreshToken) + if err != nil { + return fmt.Errorf("Your DigitalOcean session has expired and could not be renewed: %s\n\nRun `doctl auth login` to sign in again.", err) + } + + c.setContextAccessToken(token.AccessToken) + state.RefreshToken = token.RefreshToken + state.ExpiresAt = token.Expiry + if token.Scope != "" { + state.Scope = token.Scope + } + storeOAuthTokenState(authContext, state) + + return writeConfig() +} + +// accessTokenOverridden reports whether the user supplied a token for this +// invocation, in which case it takes precedence over any stored OAuth session. +func accessTokenOverridden() bool { + return Token != "" || os.Getenv("DIGITALOCEAN_ACCESS_TOKEN") != "" +} + +// currentAuthContext returns the name of the authentication context the +// command is running against. +func currentAuthContext() string { + authContext := strings.ToLower(Context) + if authContext == "" { + authContext = strings.ToLower(viper.GetString(doctl.ArgContext)) + } + if authContext == "" { + authContext = doctl.ArgDefaultContext + } + + return authContext +} + +func loadOAuthTokenState(authContext string) *oauthTokenState { + values, ok := oauthTokenStates()[strings.ToLower(authContext)].(map[string]any) + if !ok || len(values) == 0 { + return nil + } + + state := &oauthTokenState{ + Issuer: configMapString(values, "issuer"), + ClientID: configMapString(values, "client-id"), + TokenEndpoint: configMapString(values, "token-endpoint"), + RefreshToken: configMapString(values, "refresh-token"), + Scope: configMapString(values, "scope"), + } + if expiresAt := configMapString(values, "expires-at"); expiresAt != "" { + if parsed, err := time.Parse(time.RFC3339, expiresAt); err == nil { + state.ExpiresAt = parsed + } + } + if state.RefreshToken == "" { + return nil + } + + return state +} + +func storeOAuthTokenState(authContext string, state *oauthTokenState) { + var expiresAt string + if !state.ExpiresAt.IsZero() { + expiresAt = state.ExpiresAt.UTC().Format(time.RFC3339) + } + + states := oauthTokenStates() + states[strings.ToLower(authContext)] = map[string]any{ + "issuer": state.Issuer, + "client-id": state.ClientID, + "token-endpoint": state.TokenEndpoint, + "refresh-token": state.RefreshToken, + "scope": state.Scope, + "expires-at": expiresAt, + } + + viper.Set(oauthTokensConfigKey, states) +} + +// removeOAuthTokenState drops the stored OAuth session for a context. It is +// called whenever the context's credentials are replaced or removed so a stale +// refresh token cannot overwrite them later. +func removeOAuthTokenState(authContext string) { + states := oauthTokenStates() + if len(states) == 0 { + return + } + + delete(states, strings.ToLower(authContext)) + viper.Set(oauthTokensConfigKey, states) +} + +// oauthTokenStates returns a mutable copy of the stored per-context sessions, +// normalizing the values viper hands back from YAML. +func oauthTokenStates() map[string]any { + states := map[string]any{} + for name, value := range viper.GetStringMap(oauthTokensConfigKey) { + switch typed := value.(type) { + case map[string]any: + states[name] = typed + case map[string]string: + converted := make(map[string]any, len(typed)) + for k, v := range typed { + converted[k] = v + } + states[name] = converted + case map[any]any: + converted := make(map[string]any, len(typed)) + for k, v := range typed { + converted[fmt.Sprintf("%v", k)] = v + } + states[name] = converted + } + } + + return states +} + +func configMapString(values map[string]any, key string) string { + value, ok := values[key] + if !ok { + return "" + } + if s, ok := value.(string); ok { + return s + } + + return "" +} diff --git a/commands/auth_login_test.go b/commands/auth_login_test.go new file mode 100644 index 000000000..4484b4ac0 --- /dev/null +++ b/commands/auth_login_test.go @@ -0,0 +1,588 @@ +/* +Copyright 2018 The Doctl Authors All rights reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package commands + +import ( + "bytes" + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/spf13/viper" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" + + "github.com/digitalocean/doctl" + "github.com/digitalocean/doctl/do" + "github.com/digitalocean/doctl/internal/oauth" +) + +func TestRunAuthLogin(t *testing.T) { + config, token := newOAuthTestCmdConfig(t, "https://cloud.example.com") + + var capturedOpts oauth.LoginOptions + stubOAuthLogin(t, func(_ context.Context, opts oauth.LoginOptions) (*oauth.Token, error) { + capturedOpts = opts + return &oauth.Token{ + AccessToken: "doo_v1_access", + RefreshToken: "dor_v1_refresh", + Scope: "api:read api:write", + Expiry: time.Now().Add(time.Hour), + Info: oauth.TokenInfo{Email: "sammy@example.com", TeamName: "My Team"}, + }, nil + }) + + require.NoError(t, RunAuthLogin(config)) + + assert.Equal(t, "doo_v1_access", *token) + + assert.Equal(t, oauth.DefaultClientID, capturedOpts.ClientID, "doctl signs in as its own published application") + assert.Empty(t, capturedOpts.Scopes, "without --scopes the authorization screen decides the permissions") + assert.Equal(t, "https://cloud.example.com/v1/oauth/token", capturedOpts.Metadata.TokenEndpoint) + assert.False(t, capturedOpts.NoBrowser) + + state := loadOAuthTokenState(doctl.ArgDefaultContext) + require.NotNil(t, state) + assert.Equal(t, "dor_v1_refresh", state.RefreshToken) + assert.Equal(t, oauth.DefaultClientID, state.ClientID) + assert.Equal(t, "https://cloud.example.com/v1/oauth/token", state.TokenEndpoint) + assert.False(t, state.ExpiresAt.IsZero()) + assert.Equal(t, doctl.ArgDefaultContext, viper.GetString(doctl.ArgContext)) +} + +func TestRunAuthLoginWarnsThatTheCurrentContextIsReplaced(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + var out bytes.Buffer + config.Out = &out + + stubOAuthLogin(t, func(_ context.Context, _ oauth.LoginOptions) (*oauth.Token, error) { + return &oauth.Token{AccessToken: "doo_v1_access", RefreshToken: "dor_v1_refresh"}, nil + }) + + require.NoError(t, RunAuthLogin(config)) + + assert.Contains(t, out.String(), "Credentials saved there will be replaced") + assert.Contains(t, out.String(), doctl.ArgDefaultContext) + assert.Contains(t, out.String(), "doctl auth login --context ") + assert.Equal(t, doctl.ArgDefaultContext, viper.GetString(doctl.ArgContext)) +} + +func TestRunAuthLoginWithoutContextLeavesTheWorkingContextAlone(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + viper.Set(doctl.ArgContext, "your-team") + storeOAuthTokenState("your-team", &oauthTokenState{RefreshToken: "keep-me"}) + + stubOAuthLogin(t, func(_ context.Context, _ oauth.LoginOptions) (*oauth.Token, error) { + return &oauth.Token{AccessToken: "doo_v1_access", RefreshToken: "dor_v1_refresh"}, nil + }) + + require.NoError(t, RunAuthLogin(config)) + + assert.Equal(t, doctl.ArgDefaultContext, viper.GetString(doctl.ArgContext)) + require.NotNil(t, loadOAuthTokenState(doctl.ArgDefaultContext)) + assert.Equal(t, "dor_v1_refresh", loadOAuthTokenState(doctl.ArgDefaultContext).RefreshToken) + require.NotNil(t, loadOAuthTokenState("your-team")) + assert.Equal(t, "keep-me", loadOAuthTokenState("your-team").RefreshToken) +} + +func TestAuthorizationPromptIncludesTheTimeout(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + var out bytes.Buffer + config.Out = &out + + authorizationURLPrinter(config, false, 5*time.Minute)("https://example.com/authorize") + + assert.Contains(t, out.String(), "The link below expires after") + assert.Contains(t, out.String(), "5 minutes") + assert.Contains(t, out.String(), "https://example.com/authorize") +} + +func TestRunAuthLoginSwitchesToTheNamedContext(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + var out bytes.Buffer + config.Out = &out + Context = "your-team" + + stubOAuthLogin(t, func(_ context.Context, _ oauth.LoginOptions) (*oauth.Token, error) { + return &oauth.Token{AccessToken: "doo_v1_access", RefreshToken: "dor_v1_refresh"}, nil + }) + + require.NoError(t, RunAuthLogin(config)) + + assert.NotContains(t, out.String(), "Credentials saved there will be replaced") + + assert.Equal(t, "your-team", viper.GetString(doctl.ArgContext)) + require.NotNil(t, loadOAuthTokenState("your-team")) + assert.Equal(t, "dor_v1_refresh", loadOAuthTokenState("your-team").RefreshToken) + assert.Nil(t, loadOAuthTokenState(doctl.ArgDefaultContext)) +} + +func TestRunAuthLoginDoesNotCallTheAuthorizationServerBeforeTheBrowser(t *testing.T) { + var requests []string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests = append(requests, r.URL.Path) + http.NotFound(w, r) + })) + defer server.Close() + + config, _ := newOAuthTestCmdConfig(t, server.URL) + stubOAuthLogin(t, func(_ context.Context, _ oauth.LoginOptions) (*oauth.Token, error) { + return &oauth.Token{AccessToken: "doo_v1_access"}, nil + }) + + require.NoError(t, RunAuthLogin(config)) + + assert.Empty(t, requests, "endpoints are derived locally, so no metadata or registration request is made") +} + +func TestRunAuthLoginHonorsClientIDFlag(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + config.Doit.Set(config.NS, doctl.ArgOAuthClientID, "staging-client") + + var capturedOpts oauth.LoginOptions + stubOAuthLogin(t, func(_ context.Context, opts oauth.LoginOptions) (*oauth.Token, error) { + capturedOpts = opts + return &oauth.Token{AccessToken: "doo_v1_access", RefreshToken: "dor_v1_refresh"}, nil + }) + + require.NoError(t, RunAuthLogin(config)) + + assert.Equal(t, "staging-client", capturedOpts.ClientID) + assert.Equal(t, "staging-client", loadOAuthTokenState(doctl.ArgDefaultContext).ClientID) + assert.Equal(t, oauth.DefaultClientID, viper.Get(config.NS+"."+doctl.ArgOAuthClientID)) +} + +func TestRunAuthLoginClientIDAppliesOnlyToThisInvocation(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + // Left behind by an earlier `doctl auth login --client-id staging-client`, + // which writeConfig saves with the rest of the settings. + config.Doit.Set(config.NS, doctl.ArgOAuthClientID, "staging-client") + config.Doit.(*doctl.TestConfig).IsSetMap[doctl.ArgOAuthClientID] = false + + var capturedOpts oauth.LoginOptions + stubOAuthLogin(t, func(_ context.Context, opts oauth.LoginOptions) (*oauth.Token, error) { + capturedOpts = opts + return &oauth.Token{AccessToken: "doo_v1_access"}, nil + }) + + require.NoError(t, RunAuthLogin(config)) + + assert.Equal(t, oauth.DefaultClientID, capturedOpts.ClientID) +} + +func TestRunAuthLoginRejectsAnEmptyClientID(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + config.Doit.Set(config.NS, doctl.ArgOAuthClientID, "") + + err := RunAuthLogin(config) + + require.Error(t, err) + assert.Contains(t, err.Error(), "--client-id cannot be empty") +} + +func TestRunAuthLoginExplainsARejectedClient(t *testing.T) { + config, token := newOAuthTestCmdConfig(t, "https://cloud.example.com") + stubOAuthLogin(t, func(_ context.Context, _ oauth.LoginOptions) (*oauth.Token, error) { + return nil, &oauth.AuthorizationError{Code: "invalid_client"} + }) + + err := RunAuthLogin(config) + + require.Error(t, err) + assert.Contains(t, err.Error(), oauth.DefaultClientID) + assert.Empty(t, *token) +} + +func TestRunAuthLoginNoBrowserAppliesOnlyToThisInvocation(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + // Left behind by an earlier `doctl auth login --no-browser`, which + // writeConfig saves with the rest of the settings. + config.Doit.Set(config.NS, doctl.ArgOAuthNoBrowser, true) + config.Doit.(*doctl.TestConfig).IsSetMap[doctl.ArgOAuthNoBrowser] = false + + var capturedOpts oauth.LoginOptions + stubOAuthLogin(t, func(_ context.Context, opts oauth.LoginOptions) (*oauth.Token, error) { + capturedOpts = opts + return &oauth.Token{AccessToken: "doo_v1_access", RefreshToken: "dor_v1_refresh"}, nil + }) + + require.NoError(t, RunAuthLogin(config)) + assert.False(t, capturedOpts.NoBrowser, "a saved --no-browser must not keep the browser closed") + assert.Equal(t, false, viper.Get(config.NS+"."+doctl.ArgOAuthNoBrowser)) +} + +func TestRunAuthLoginHonorsNoBrowserFlag(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + config.Doit.Set(config.NS, doctl.ArgOAuthNoBrowser, true) + + var capturedOpts oauth.LoginOptions + stubOAuthLogin(t, func(_ context.Context, opts oauth.LoginOptions) (*oauth.Token, error) { + capturedOpts = opts + return &oauth.Token{AccessToken: "doo_v1_access", RefreshToken: "dor_v1_refresh"}, nil + }) + + require.NoError(t, RunAuthLogin(config)) + assert.True(t, capturedOpts.NoBrowser) +} + +func TestRunAuthLoginRequestsTheGivenScopes(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + config.Doit.Set(config.NS, doctl.ArgOAuthScopes, "read write") + + var capturedOpts oauth.LoginOptions + stubOAuthLogin(t, func(_ context.Context, opts oauth.LoginOptions) (*oauth.Token, error) { + capturedOpts = opts + return &oauth.Token{AccessToken: "doo_v1_access"}, nil + }) + + require.NoError(t, RunAuthLogin(config)) + + assert.Equal(t, []string{"read", "write"}, capturedOpts.Scopes) + assert.Equal(t, defaultOAuthScopes, viper.Get(config.NS+"."+doctl.ArgOAuthScopes)) +} + +func TestRunAuthLoginScopeAppliesOnlyToThisInvocation(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + // Left behind by an earlier `doctl auth login --scopes "read write"`, which + // writeConfig saves with the rest of the settings. + config.Doit.Set(config.NS, doctl.ArgOAuthScopes, "read write") + config.Doit.(*doctl.TestConfig).IsSetMap[doctl.ArgOAuthScopes] = false + + var capturedOpts oauth.LoginOptions + stubOAuthLogin(t, func(_ context.Context, opts oauth.LoginOptions) (*oauth.Token, error) { + capturedOpts = opts + return &oauth.Token{AccessToken: "doo_v1_access", RefreshToken: "dor_v1_refresh"}, nil + }) + + require.NoError(t, RunAuthLogin(config)) + assert.Empty(t, capturedOpts.Scopes, "a saved --scopes must not be reused on later logins") + assert.Equal(t, defaultOAuthScopes, viper.Get(config.NS+"."+doctl.ArgOAuthScopes)) +} + +func TestRunAuthLoginSaveScopePersistsDefault(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + config.Doit.Set(config.NS, doctl.ArgOAuthScopes, "read write") + config.Doit.Set(config.NS, doctl.ArgOAuthSaveScope, true) + + stubOAuthLogin(t, func(_ context.Context, opts oauth.LoginOptions) (*oauth.Token, error) { + assert.Equal(t, []string{"read", "write"}, opts.Scopes) + return &oauth.Token{AccessToken: "doo_v1_access", RefreshToken: "dor_v1_refresh"}, nil + }) + + require.NoError(t, RunAuthLogin(config)) + assert.Equal(t, "read write", loadOAuthDefaultScopes()) + assert.Equal(t, false, viper.Get(config.NS+"."+doctl.ArgOAuthSaveScope)) +} + +func TestRunAuthLoginUsesSavedDefaultScopes(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + storeOAuthDefaultScopes("droplet:read account:read") + + var capturedOpts oauth.LoginOptions + stubOAuthLogin(t, func(_ context.Context, opts oauth.LoginOptions) (*oauth.Token, error) { + capturedOpts = opts + return &oauth.Token{AccessToken: "doo_v1_access", RefreshToken: "dor_v1_refresh"}, nil + }) + + require.NoError(t, RunAuthLogin(config)) + assert.Equal(t, []string{"droplet:read", "account:read"}, capturedOpts.Scopes) +} + +func TestRunAuthLoginExplicitScopeOverridesSavedDefault(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + storeOAuthDefaultScopes("read write") + config.Doit.Set(config.NS, doctl.ArgOAuthScopes, "droplet:read") + + var capturedOpts oauth.LoginOptions + stubOAuthLogin(t, func(_ context.Context, opts oauth.LoginOptions) (*oauth.Token, error) { + capturedOpts = opts + return &oauth.Token{AccessToken: "doo_v1_access"}, nil + }) + + require.NoError(t, RunAuthLogin(config)) + assert.Equal(t, []string{"droplet:read"}, capturedOpts.Scopes) + assert.Equal(t, "read write", loadOAuthDefaultScopes(), "override without --save-scope must leave the default alone") +} + +func TestRunAuthLoginSaveScopeClearsDefault(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + storeOAuthDefaultScopes("read write") + config.Doit.Set(config.NS, doctl.ArgOAuthScopes, "") + config.Doit.Set(config.NS, doctl.ArgOAuthSaveScope, true) + + stubOAuthLogin(t, func(_ context.Context, _ oauth.LoginOptions) (*oauth.Token, error) { + return &oauth.Token{AccessToken: "doo_v1_access"}, nil + }) + + require.NoError(t, RunAuthLogin(config)) + assert.Empty(t, loadOAuthDefaultScopes()) +} + +func TestRunAuthLoginSaveScopeRequiresScopeFlag(t *testing.T) { + config, _ := newOAuthTestCmdConfig(t, "https://cloud.example.com") + config.Doit.Set(config.NS, doctl.ArgOAuthSaveScope, true) + + err := RunAuthLogin(config) + require.Error(t, err) + assert.Contains(t, err.Error(), "--save-scope requires --scopes") +} + +func TestRunAuthLoginFailure(t *testing.T) { + config, token := newOAuthTestCmdConfig(t, "https://cloud.example.com") + stubOAuthLogin(t, func(_ context.Context, opts oauth.LoginOptions) (*oauth.Token, error) { + return nil, &oauth.AuthorizationError{Code: "access_denied"} + }) + + err := RunAuthLogin(config) + + require.Error(t, err) + assert.Contains(t, err.Error(), "access_denied") + assert.Empty(t, *token, "a failed login must not store a token") + assert.Nil(t, loadOAuthTokenState(doctl.ArgDefaultContext)) +} + +func TestRefreshExpiredOAuthToken(t *testing.T) { + tokenEndpoint := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.NoError(t, r.ParseForm()) + assert.Equal(t, "refresh_token", r.FormValue("grant_type")) + assert.Equal(t, "dor_v1_old", r.FormValue("refresh_token")) + + w.Header().Set("Content-Type", "application/json") + require.NoError(t, json.NewEncoder(w).Encode(map[string]any{ + "access_token": "doo_v1_new", + "refresh_token": "dor_v1_new", + "expires_in": 3600, + })) + })) + defer tokenEndpoint.Close() + + t.Run("renews an expired token", func(t *testing.T) { + config, token := newOAuthTestCmdConfig(t, tokenEndpoint.URL) + storeOAuthTokenState(doctl.ArgDefaultContext, &oauthTokenState{ + Issuer: tokenEndpoint.URL, + ClientID: "client-1", + TokenEndpoint: tokenEndpoint.URL, + RefreshToken: "dor_v1_old", + ExpiresAt: time.Now().Add(-time.Minute), + }) + + require.NoError(t, refreshExpiredOAuthToken(config)) + + assert.Equal(t, "doo_v1_new", *token) + + state := loadOAuthTokenState(doctl.ArgDefaultContext) + require.NotNil(t, state) + assert.Equal(t, "dor_v1_new", state.RefreshToken) + assert.True(t, state.ExpiresAt.After(time.Now())) + }) + + t.Run("leaves a valid token alone", func(t *testing.T) { + config, token := newOAuthTestCmdConfig(t, tokenEndpoint.URL) + storeOAuthTokenState(doctl.ArgDefaultContext, &oauthTokenState{ + TokenEndpoint: tokenEndpoint.URL, + RefreshToken: "dor_v1_old", + ExpiresAt: time.Now().Add(time.Hour), + }) + + require.NoError(t, refreshExpiredOAuthToken(config)) + + assert.Empty(t, *token) + }) + + t.Run("ignores contexts without an OAuth session", func(t *testing.T) { + config, token := newOAuthTestCmdConfig(t, tokenEndpoint.URL) + + require.NoError(t, refreshExpiredOAuthToken(config)) + + assert.Empty(t, *token) + }) + + t.Run("ignores a token supplied for this command", func(t *testing.T) { + config, token := newOAuthTestCmdConfig(t, tokenEndpoint.URL) + storeOAuthTokenState(doctl.ArgDefaultContext, &oauthTokenState{ + TokenEndpoint: tokenEndpoint.URL, + RefreshToken: "dor_v1_old", + ExpiresAt: time.Now().Add(-time.Minute), + }) + + Token = "dop_v1_explicit" + t.Cleanup(func() { Token = "" }) + + require.NoError(t, refreshExpiredOAuthToken(config)) + + assert.Empty(t, *token) + }) + + t.Run("explains how to recover when the refresh token is rejected", func(t *testing.T) { + rejecting := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusUnauthorized) + require.NoError(t, json.NewEncoder(w).Encode(map[string]any{"error": "invalid_grant"})) + })) + defer rejecting.Close() + + config, _ := newOAuthTestCmdConfig(t, rejecting.URL) + storeOAuthTokenState(doctl.ArgDefaultContext, &oauthTokenState{ + TokenEndpoint: rejecting.URL, + RefreshToken: "dor_v1_old", + ExpiresAt: time.Now().Add(-time.Minute), + }) + + err := refreshExpiredOAuthToken(config) + + require.Error(t, err) + assert.Contains(t, err.Error(), "doctl auth login") + }) +} + +func TestAuthInitDiscardsTheOAuthSession(t *testing.T) { + resetOAuthConfig(t) + storeOAuthTokenState(doctl.ArgDefaultContext, &oauthTokenState{ + RefreshToken: "dor_v1_refresh", + ExpiresAt: time.Now().Add(-time.Minute), + }) + + previousToken := viper.Get(doctl.ArgAccessToken) + viper.Set(doctl.ArgAccessToken, "dop_v1_pat") + t.Cleanup(func() { viper.Set(doctl.ArgAccessToken, previousToken) }) + + withTestClient(t, func(config *CmdConfig, tm *tcMocks) { + tm.oauth.EXPECT().TokenInfo(gomock.Any()).Return(&do.OAuthTokenInfo{}, nil) + + err := RunAuthInit(func() (string, error) { return "", nil })(config) + require.NoError(t, err) + }) + + assert.Nil(t, loadOAuthTokenState(doctl.ArgDefaultContext)) +} + +func TestOAuthTokenStateIsScopedToAContext(t *testing.T) { + resetOAuthConfig(t) + + storeOAuthTokenState("default", &oauthTokenState{RefreshToken: "default-refresh", Scope: "read"}) + storeOAuthTokenState("Your-Team", &oauthTokenState{RefreshToken: "team-refresh"}) + + assert.Equal(t, "default-refresh", loadOAuthTokenState("default").RefreshToken) + assert.Equal(t, "team-refresh", loadOAuthTokenState("your-team").RefreshToken, "context names are case insensitive") + assert.Nil(t, loadOAuthTokenState("unknown")) + + removeOAuthTokenState("default") + + assert.Nil(t, loadOAuthTokenState("default")) + assert.NotNil(t, loadOAuthTokenState("your-team")) +} + +func TestOAuthTokenStateNeedsRefresh(t *testing.T) { + now := time.Now() + + tests := []struct { + name string + state *oauthTokenState + expected bool + }{ + {name: "no session"}, + {name: "no refresh token", state: &oauthTokenState{ExpiresAt: now.Add(-time.Hour)}}, + {name: "unknown expiry", state: &oauthTokenState{RefreshToken: "refresh"}}, + {name: "valid", state: &oauthTokenState{RefreshToken: "refresh", ExpiresAt: now.Add(time.Hour)}, expected: false}, + {name: "expired", state: &oauthTokenState{RefreshToken: "refresh", ExpiresAt: now.Add(-time.Second)}, expected: true}, + {name: "expiring within the skew", state: &oauthTokenState{RefreshToken: "refresh", ExpiresAt: now.Add(time.Minute)}, expected: true}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + assert.Equal(t, test.expected, test.state.needsRefresh(now)) + }) + } +} + +// newOAuthTestCmdConfig returns a command config wired to an isolated view of +// the doctl configuration, along with a pointer to the access token the +// command stores for the current context. +func newOAuthTestCmdConfig(t *testing.T, issuer string) (*CmdConfig, *string) { + t.Helper() + + resetOAuthConfig(t) + + doitConfig := doctl.NewTestConfig() + var storedToken string + + config := &CmdConfig{ + NS: "test", + Doit: doitConfig, + Out: io.Discard, + setContextAccessToken: func(token string) { storedToken = token }, + getContextAccessToken: func() string { return storedToken }, + initServices: func(*CmdConfig) error { return nil }, + } + + doitConfig.Set(config.NS, doctl.ArgOAuthServer, issuer) + doitConfig.Set(config.NS, doctl.ArgOAuthTimeout, time.Minute) + + return config, &storedToken +} + +// resetOAuthConfig isolates a test from the global viper configuration the +// commands package writes to. +func resetOAuthConfig(t *testing.T) { + t.Helper() + + previousWriter := cfgFileWriter + previousContext := Context + previousTokens := viper.Get(oauthTokensConfigKey) + previousDefaultScopes := viper.Get(oauthDefaultScopesConfigKey) + previousViperContext := viper.Get(doctl.ArgContext) + ephemeralKeys := []string{ + "test." + doctl.ArgOAuthServer, + "test." + doctl.ArgOAuthClientID, + "test." + doctl.ArgOAuthScopes, + "test." + doctl.ArgOAuthSaveScope, + "test." + doctl.ArgOAuthCallbackPort, + "test." + doctl.ArgOAuthNoBrowser, + "test." + doctl.ArgOAuthTimeout, + } + previousEphemeral := make(map[string]any, len(ephemeralKeys)) + for _, key := range ephemeralKeys { + previousEphemeral[key] = viper.Get(key) + } + + cfgFileWriter = func() (io.WriteCloser, error) { return &nopWriteCloser{Writer: io.Discard}, nil } + Context = "" + viper.Set(doctl.ArgContext, doctl.ArgDefaultContext) + viper.Set(oauthTokensConfigKey, nil) + viper.Set(oauthDefaultScopesConfigKey, nil) + + t.Cleanup(func() { + cfgFileWriter = previousWriter + Context = previousContext + viper.Set(oauthTokensConfigKey, previousTokens) + viper.Set(oauthDefaultScopesConfigKey, previousDefaultScopes) + viper.Set(doctl.ArgContext, previousViperContext) + for key, value := range previousEphemeral { + viper.Set(key, value) + } + }) +} + +func stubOAuthLogin(t *testing.T, login func(context.Context, oauth.LoginOptions) (*oauth.Token, error)) { + t.Helper() + + original := oauth.Login + oauth.Login = login + t.Cleanup(func() { oauth.Login = original }) +} diff --git a/commands/auth_test.go b/commands/auth_test.go index 4a30f014a..1b3f63f05 100644 --- a/commands/auth_test.go +++ b/commands/auth_test.go @@ -32,7 +32,7 @@ import ( func TestAuthCommand(t *testing.T) { cmd := Auth() assert.NotNil(t, cmd) - assertCommandNames(t, cmd, "init", "list", "remove", "switch", "token") + assertCommandNames(t, cmd, "init", "list", "login", "remove", "switch", "token") } func TestAuthInit(t *testing.T) { diff --git a/commands/command_config.go b/commands/command_config.go index 26ad08638..219a90966 100644 --- a/commands/command_config.go +++ b/commands/command_config.go @@ -108,6 +108,10 @@ func NewCmdConfig(ns string, dc doctl.Config, out io.Writer, args []string, init Args: args, initServices: func(c *CmdConfig) error { + if err := refreshExpiredOAuthToken(c); err != nil { + return err + } + accessToken := c.getContextAccessToken() godoClient, err := c.Doit.GetGodoClient(Trace, true, accessToken) if err != nil { diff --git a/integration/auth_login_test.go b/integration/auth_login_test.go new file mode 100644 index 000000000..c58c295c3 --- /dev/null +++ b/integration/auth_login_test.go @@ -0,0 +1,537 @@ +package integration + +import ( + "bytes" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "net/url" + "os" + "os/exec" + "path/filepath" + "regexp" + "strings" + "sync" + "testing" + "time" + + "github.com/sclevine/spec" + "github.com/stretchr/testify/require" +) + +// ansiEscape matches the styling the CLI applies to the authorization URL when +// it renders it, so the test can recover the bare URL. +var ansiEscape = regexp.MustCompile(`\x1b\[[0-9;]*[a-zA-Z]`) + +// authorizationServer is a stand-in for DigitalOcean's authorization server. +// It records what doctl sent to the authorize endpoint so the token endpoint +// can verify the PKCE verifier against the challenge from the same login. +type authorizationServer struct { + *httptest.Server + + mu sync.Mutex + codeChallenge string + challengeKind string + authClientID string + authScope string + redirectURI string + tokenForm url.Values +} + +func newAuthorizationServer(t *testing.T) *authorizationServer { + as := &authorizationServer{} + + mux := http.NewServeMux() + + // The browser lands here. A real authorization server would render a + // consent screen; this one approves immediately and redirects back to the + // loopback address doctl is listening on. + mux.HandleFunc("/v1/oauth/authorize", func(w http.ResponseWriter, r *http.Request) { + query := r.URL.Query() + + as.mu.Lock() + as.codeChallenge = query.Get("code_challenge") + as.challengeKind = query.Get("code_challenge_method") + as.authClientID = query.Get("client_id") + as.authScope = query.Get("scope") + as.redirectURI = query.Get("redirect_uri") + as.mu.Unlock() + + redirect, err := url.Parse(query.Get("redirect_uri")) + if err != nil { + http.Error(w, "bad redirect_uri", http.StatusBadRequest) + return + } + + params := redirect.Query() + params.Set("code", "the-authorization-code") + params.Set("state", query.Get("state")) + redirect.RawQuery = params.Encode() + + http.Redirect(w, r, redirect.String(), http.StatusFound) + }) + + mux.HandleFunc("/v1/oauth/token", func(w http.ResponseWriter, r *http.Request) { + if err := r.ParseForm(); err != nil { + http.Error(w, "bad form", http.StatusBadRequest) + return + } + + as.mu.Lock() + challenge := as.codeChallenge + as.tokenForm = r.PostForm + as.mu.Unlock() + + // Proving the verifier against the challenge is the point of PKCE, so + // reject the exchange rather than hand out a token when it fails. + sum := sha256.Sum256([]byte(r.FormValue("code_verifier"))) + if base64.RawURLEncoding.EncodeToString(sum[:]) != challenge { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + json.NewEncoder(w).Encode(map[string]any{ + "error": "invalid_grant", + "error_description": "the code verifier did not match the code challenge", + }) + return + } + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "access_token": "doo_v1_from_the_code_exchange", + "refresh_token": "dor_v1_from_the_code_exchange", + "token_type": "Bearer", + "expires_in": 3600, + "scope": "read write", + "info": map[string]any{ + "email": "sammy@example.com", + "team_name": "My Team", + }, + }) + }) + + as.Server = httptest.NewServer(mux) + t.Cleanup(as.Close) + + return as +} + +func (a *authorizationServer) snapshot() (challenge, kind, clientID, scope, redirectURI string, tokenForm url.Values) { + a.mu.Lock() + defer a.mu.Unlock() + + return a.codeChallenge, a.challengeKind, a.authClientID, a.authScope, a.redirectURI, a.tokenForm +} + +var _ = suite("auth/login", func(t *testing.T, when spec.G, it spec.S) { + var expect *require.Assertions + + it.Before(func() { + expect = require.New(t) + }) + + when("the browser completes the authorization", func() { + it("exchanges the code and saves the session", func() { + as := newAuthorizationServer(t) + + tmpDir := t.TempDir() + testConfig := filepath.Join(tmpDir, "test-config.yml") + + cmd := exec.Command(builtBinaryPath, + "--config", testConfig, + "auth", "login", + "--oauth-server", as.URL, + "--client-id", "doctl-e2e-client", + "--scopes", "read write", + "--no-browser", + "--timeout", "30s", + ) + + output, authURL := startLoginAndCaptureURL(t, expect, cmd) + + // Stand in for the browser: follow the authorization URL, which + // redirects back into doctl's loopback callback server. + visitAuthorizationURL(t, expect, authURL) + + expect.NoError(cmd.Wait(), output.String()) + + combined := output.String() + expect.Contains(combined, "Signed in as") + expect.Contains(combined, "sammy@example.com") + expect.Contains(combined, "My Team") + + challenge, kind, clientID, scope, redirectURI, tokenForm := as.snapshot() + + expect.Equal("doctl-e2e-client", clientID, "doctl must authorize as the configured application") + expect.Equal("read write", scope) + expect.Equal("S256", kind, "doctl must use the S256 PKCE challenge method") + expect.NotEmpty(challenge) + expect.Regexp(`^http://127\.0\.0\.1:\d+/callback$`, redirectURI) + + expect.Equal("authorization_code", tokenForm.Get("grant_type")) + expect.Equal("the-authorization-code", tokenForm.Get("code")) + expect.Equal("doctl-e2e-client", tokenForm.Get("client_id")) + expect.Equal(redirectURI, tokenForm.Get("redirect_uri")) + expect.Empty(tokenForm.Get("client_secret"), "doctl is a public client and must not send a secret") + + config := readConfigFile(t, expect, testConfig) + expect.Contains(config, "access-token: doo_v1_from_the_code_exchange") + expect.Contains(config, "refresh-token: dor_v1_from_the_code_exchange") + expect.NotContains(config, "the-authorization-code", "the authorization code is single use and must not be stored") + }) + }) + + when("the authorization server reports an error", func() { + it("exits non-zero without saving a token", func() { + denying := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + query := r.URL.Query() + + redirect, err := url.Parse(query.Get("redirect_uri")) + if err != nil { + http.Error(w, "bad redirect_uri", http.StatusBadRequest) + return + } + + params := redirect.Query() + params.Set("error", "access_denied") + params.Set("error_description", "the user declined the request") + params.Set("state", query.Get("state")) + redirect.RawQuery = params.Encode() + + http.Redirect(w, r, redirect.String(), http.StatusFound) + })) + defer denying.Close() + + tmpDir := t.TempDir() + testConfig := filepath.Join(tmpDir, "test-config.yml") + + cmd := exec.Command(builtBinaryPath, + "--config", testConfig, + "auth", "login", + "--oauth-server", denying.URL, + "--client-id", "doctl-e2e-client", + "--no-browser", + "--timeout", "30s", + ) + + output, authURL := startLoginAndCaptureURL(t, expect, cmd) + visitAuthorizationURL(t, expect, authURL) + + err := cmd.Wait() + expect.Error(err, "a declined authorization must fail the command") + + combined := output.String() + expect.Contains(combined, "access_denied") + expect.Contains(combined, "the user declined the request") + + if contents, err := os.ReadFile(testConfig); err == nil { + expect.NotContains(string(contents), "access-token: doo_v1") + } + }) + }) +}) + +var _ = suite("auth/login/refresh", func(t *testing.T, when spec.G, it spec.S) { + var expect *require.Assertions + + it.Before(func() { + expect = require.New(t) + }) + + when("the stored access token has expired", func() { + it("renews it before running the command and saves the rotated tokens", func() { + var ( + mu sync.Mutex + refreshForm url.Values + accountAuth string + ) + + mux := http.NewServeMux() + + mux.HandleFunc("/v1/oauth/token", func(w http.ResponseWriter, r *http.Request) { + expect.NoError(r.ParseForm()) + + mu.Lock() + refreshForm = r.PostForm + mu.Unlock() + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "access_token": "doo_v1_renewed", + "refresh_token": "dor_v1_rotated", + "token_type": "Bearer", + "expires_in": 3600, + "scope": "read write", + }) + }) + + mux.HandleFunc("/v2/account", func(w http.ResponseWriter, r *http.Request) { + mu.Lock() + accountAuth = r.Header.Get("Authorization") + mu.Unlock() + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "account": map[string]any{ + "email": "sammy@example.com", + "uuid": "a-uuid", + "status": "active", + "droplet_limit": 25, + "email_verified": true, + }, + }) + }) + + server := httptest.NewServer(mux) + defer server.Close() + + tmpDir := t.TempDir() + testConfig := filepath.Join(tmpDir, "test-config.yml") + + expired := time.Now().Add(-time.Hour).UTC().Format(time.RFC3339) + expect.NoError(os.WriteFile(testConfig, []byte(strings.Join([]string{ + "access-token: doo_v1_expired", + "context: default", + "oauth-tokens:", + " default:", + " issuer: " + server.URL, + " client-id: doctl-e2e-client", + " token-endpoint: " + server.URL + "/v1/oauth/token", + " refresh-token: dor_v1_stored", + " scope: read write", + " expires-at: " + expired, + "", + }, "\n")), 0600)) + + cmd := exec.Command(builtBinaryPath, + "-u", server.URL, + "--config", testConfig, + "account", "get", + ) + cmd.Env = environmentWithoutDigitalOceanAuth() + + output, err := cmd.CombinedOutput() + expect.NoError(err, string(output)) + expect.Contains(string(output), "sammy@example.com") + + mu.Lock() + form, auth := refreshForm, accountAuth + mu.Unlock() + + expect.NotNil(form, "doctl must renew the token before calling the API") + expect.Equal("refresh_token", form.Get("grant_type")) + expect.Equal("dor_v1_stored", form.Get("refresh_token")) + expect.Equal("doctl-e2e-client", form.Get("client_id")) + + expect.Equal("Bearer doo_v1_renewed", auth, "the API call must use the renewed token") + + config := readConfigFile(t, expect, testConfig) + expect.Contains(config, "access-token: doo_v1_renewed") + expect.Contains(config, "refresh-token: dor_v1_rotated", "a rotated refresh token must replace the stored one") + expect.NotContains(config, "dor_v1_stored", "the consumed refresh token must not be left behind") + }) + }) + + when("the stored access token is still valid", func() { + it("uses it without contacting the token endpoint", func() { + var ( + mu sync.Mutex + refreshed bool + auth string + ) + + mux := http.NewServeMux() + + mux.HandleFunc("/v1/oauth/token", func(w http.ResponseWriter, r *http.Request) { + mu.Lock() + refreshed = true + mu.Unlock() + + w.WriteHeader(http.StatusInternalServerError) + }) + + mux.HandleFunc("/v2/account", func(w http.ResponseWriter, r *http.Request) { + mu.Lock() + auth = r.Header.Get("Authorization") + mu.Unlock() + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "account": map[string]any{ + "email": "sammy@example.com", + "uuid": "a-uuid", + "status": "active", + }, + }) + }) + + server := httptest.NewServer(mux) + defer server.Close() + + tmpDir := t.TempDir() + testConfig := filepath.Join(tmpDir, "test-config.yml") + + valid := time.Now().Add(time.Hour).UTC().Format(time.RFC3339) + expect.NoError(os.WriteFile(testConfig, []byte(strings.Join([]string{ + "access-token: doo_v1_still_good", + "context: default", + "oauth-tokens:", + " default:", + " issuer: " + server.URL, + " client-id: doctl-e2e-client", + " token-endpoint: " + server.URL + "/v1/oauth/token", + " refresh-token: dor_v1_stored", + " scope: read write", + " expires-at: " + valid, + "", + }, "\n")), 0600)) + + cmd := exec.Command(builtBinaryPath, + "-u", server.URL, + "--config", testConfig, + "account", "get", + ) + cmd.Env = environmentWithoutDigitalOceanAuth() + + output, err := cmd.CombinedOutput() + expect.NoError(err, string(output)) + + mu.Lock() + didRefresh, usedAuth := refreshed, auth + mu.Unlock() + + expect.False(didRefresh, "an unexpired token must not be renewed") + expect.Equal("Bearer doo_v1_still_good", usedAuth) + }) + }) +}) + +// environmentWithoutDigitalOceanAuth drops the ambient credentials a developer +// may have exported, which would otherwise take precedence over the config +// file the test wrote. +func environmentWithoutDigitalOceanAuth() []string { + var env []string + for _, entry := range os.Environ() { + switch { + case strings.HasPrefix(entry, "DIGITALOCEAN_ACCESS_TOKEN="), + strings.HasPrefix(entry, "DIGITALOCEAN_CONTEXT="): + continue + } + env = append(env, entry) + } + + return env +} + +// startLoginAndCaptureURL runs the login command and returns the captured +// output along with the authorization URL doctl printed for the browser. The +// command is still running when this returns: it is waiting on its callback. +// +// Output is collected in the writer the command prints to, rather than from a +// pipe read after Wait. Wait closes a StdoutPipe, which can discard the error +// line doctl writes as it exits. +func startLoginAndCaptureURL(t *testing.T, expect *require.Assertions, cmd *exec.Cmd) (*safeBuffer, string) { + t.Helper() + + output := &safeBuffer{} + catcher := &urlCatcher{buf: output, urls: make(chan string, 1)} + cmd.Stdout = catcher + cmd.Stderr = catcher + + expect.NoError(cmd.Start()) + + select { + case authURL := <-catcher.urls: + return output, authURL + case <-time.After(30 * time.Second): + _ = cmd.Process.Kill() + t.Fatalf("timed out waiting for the authorization URL: %s", output.String()) + return output, "" + } +} + +// urlCatcher records command output and reports the first line that is an +// authorization URL. Writes are complete before the process exits, so the +// recorded output is complete once Wait returns. +type urlCatcher struct { + mu sync.Mutex + buf *safeBuffer + urls chan string + rest []byte +} + +func (c *urlCatcher) Write(p []byte) (int, error) { + c.mu.Lock() + defer c.mu.Unlock() + + if _, err := c.buf.Write(p); err != nil { + return 0, err + } + + c.rest = append(c.rest, p...) + for { + newline := bytes.IndexByte(c.rest, '\n') + if newline < 0 { + break + } + line := strings.TrimSpace(ansiEscape.ReplaceAllString(string(c.rest[:newline]), "")) + c.rest = c.rest[newline+1:] + if strings.HasPrefix(line, "http://") || strings.HasPrefix(line, "https://") { + select { + case c.urls <- line: + default: + } + } + } + + return len(p), nil +} + +// visitAuthorizationURL acts as the user's browser. The authorization server +// redirects to doctl's loopback callback, which the client follows. +func visitAuthorizationURL(t *testing.T, expect *require.Assertions, authURL string) { + t.Helper() + + resp, err := http.Get(authURL) + expect.NoError(err) + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + expect.NoError(err) + + // Whatever the outcome, the callback page is what the user is left looking + // at, so it must not be a bare error from the Go http package. + expect.Contains(string(body), "") +} + +func readConfigFile(t *testing.T, expect *require.Assertions, path string) string { + t.Helper() + + contents, err := os.ReadFile(path) + expect.NoError(err) + + return string(contents) +} + +// safeBuffer collects command output that is written from the scanning +// goroutine and read from the test goroutine. +type safeBuffer struct { + mu sync.Mutex + builder strings.Builder +} + +func (b *safeBuffer) Write(p []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + + return b.builder.Write(p) +} + +func (b *safeBuffer) String() string { + b.mu.Lock() + defer b.mu.Unlock() + + return b.builder.String() +} diff --git a/internal/oauth/login.go b/internal/oauth/login.go new file mode 100644 index 000000000..dec46376a --- /dev/null +++ b/internal/oauth/login.go @@ -0,0 +1,270 @@ +/* +Copyright 2018 The Doctl Authors All rights reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package oauth + +import ( + "context" + "crypto/rand" + "crypto/sha256" + "crypto/subtle" + "encoding/base64" + "errors" + "fmt" + "net" + "net/http" + "net/url" + "strings" + "time" + + "github.com/pkg/browser" +) + +// callbackPath is the path the local redirect server listens on. It is part of +// the registered redirect URIs, so it must stay stable across releases. +const callbackPath = "/callback" + +// DefaultLoginTimeout bounds how long the login flow waits for the user to +// finish authorizing in their browser. +const DefaultLoginTimeout = 5 * time.Minute + +// RedirectURIs are the redirect URIs the doctl OAuth application registers +// with DigitalOcean. Neither carries a port: the authorization server matches +// loopback redirects on any port (RFC 8252 section 7.3), which lets doctl bind +// an ephemeral port at login time. +func RedirectURIs() []string { + return []string{ + "http://127.0.0.1" + callbackPath, + "http://localhost" + callbackPath, + } +} + +// AuthorizationError is an error the authorization server reported by +// redirecting back to doctl with an error parameter. +type AuthorizationError struct { + Code string + Description string +} + +func (e *AuthorizationError) Error() string { + if e.Description != "" { + return fmt.Sprintf("%s: %s", e.Code, e.Description) + } + return e.Code +} + +// Unwrap reports a rejected client as ErrInvalidClient so callers can explain +// that the doctl application, rather than the user, was turned away. +func (e *AuthorizationError) Unwrap() error { + if isInvalidClientCode(e.Code) { + return ErrInvalidClient + } + return nil +} + +// LoginOptions configures a single authorization code login. +type LoginOptions struct { + // Metadata describes the authorization server endpoints to use. + Metadata *ServerMetadata + // ClientID identifies the dynamically registered public client. + ClientID string + // Scopes are the OAuth scopes to request. An empty list lets the + // authorization server apply its default scope. + Scopes []string + // Port is the local port to listen on for the redirect. Zero picks an + // unused port, which is the recommended behavior for native apps. + Port int + // Timeout bounds how long to wait for the browser redirect. Defaults to + // DefaultLoginTimeout. + Timeout time.Duration + // HTTPClient is used for the token exchange. + HTTPClient *http.Client + // NoBrowser prints the authorization URL instead of opening a browser. + NoBrowser bool + // OnAuthorizationURL, when set, is called with the authorization URL once + // the local redirect server is listening. + OnAuthorizationURL func(authURL string) + + // openURL is a test hook for opening the browser. + openURL func(authURL string) error +} + +// Login runs the OAuth 2.1 authorization code flow with PKCE against the +// configured authorization server and returns the resulting token grant. +// It is a variable so tests can run the command without opening a browser. +var Login = login + +func login(ctx context.Context, opts LoginOptions) (*Token, error) { + if opts.Metadata == nil { + return nil, errors.New("authorization server metadata is required") + } + if opts.ClientID == "" { + return nil, errors.New("a client ID is required") + } + if opts.Timeout <= 0 { + opts.Timeout = DefaultLoginTimeout + } + if opts.openURL == nil { + opts.openURL = browser.OpenURL + } + + verifier, err := generateCodeVerifier() + if err != nil { + return nil, err + } + state, err := randomURLSafeString(32) + if err != nil { + return nil, err + } + + listener, err := net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", opts.Port)) + if err != nil { + return nil, fmt.Errorf("listening for the authorization redirect: %w", err) + } + defer listener.Close() + + redirectURI := fmt.Sprintf("http://127.0.0.1:%d%s", listener.Addr().(*net.TCPAddr).Port, callbackPath) + authURL := authorizationURL(opts, redirectURI, state, verifier) + + ctx, cancel := context.WithTimeout(ctx, opts.Timeout) + defer cancel() + + results := make(chan callbackResult, 1) + server := &http.Server{Handler: &callbackHandler{state: state, results: results}} + defer server.Close() + + go func() { _ = server.Serve(listener) }() + + if opts.OnAuthorizationURL != nil { + opts.OnAuthorizationURL(authURL) + } + if !opts.NoBrowser { + if err := opts.openURL(authURL); err != nil { + return nil, fmt.Errorf("opening the authorization URL in a browser: %w", err) + } + } + + var result callbackResult + select { + case result = <-results: + case <-ctx.Done(): + if errors.Is(ctx.Err(), context.DeadlineExceeded) { + return nil, fmt.Errorf("timed out after %s waiting for browser authorization", opts.Timeout) + } + return nil, ctx.Err() + } + + if result.err != nil { + return nil, result.err + } + + return ExchangeCode(ctx, opts.HTTPClient, opts.Metadata.TokenEndpoint, opts.ClientID, result.code, verifier, redirectURI) +} + +func authorizationURL(opts LoginOptions, redirectURI, state, verifier string) string { + query := url.Values{ + "client_id": {opts.ClientID}, + "redirect_uri": {redirectURI}, + "response_type": {responseTypeCode}, + "state": {state}, + "code_challenge": {codeChallengeS256(verifier)}, + "code_challenge_method": {CodeChallengeMethodS256}, + } + if scope := strings.TrimSpace(strings.Join(opts.Scopes, " ")); scope != "" { + query.Set("scope", scope) + } + + separator := "?" + if strings.Contains(opts.Metadata.AuthorizationEndpoint, "?") { + separator = "&" + } + + return opts.Metadata.AuthorizationEndpoint + separator + query.Encode() +} + +type callbackResult struct { + code string + err error +} + +type callbackHandler struct { + state string + results chan<- callbackResult +} + +func (h *callbackHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != callbackPath { + http.NotFound(w, r) + return + } + + query := r.URL.Query() + result, content, status := h.resultFor(query) + + w.Header().Set("Content-Type", "text/html; charset=utf-8") + w.WriteHeader(status) + // Execute the html/template onto the response so text from the + // authorization server is escaped for the HTML context it appears in. + // Rendering to a string and writing those bytes would hide that escaping + // from the response path. + if err := pageTemplate.Execute(w, content); err != nil { + _, _ = w.Write([]byte("Return to your terminal to continue.")) + } + + select { + case h.results <- result: + default: + } +} + +func (h *callbackHandler) resultFor(query url.Values) (callbackResult, pageContent, int) { + // The state check comes first: a mismatched state means the redirect did + // not originate from the authorization request we started. + if subtle.ConstantTimeCompare([]byte(query.Get("state")), []byte(h.state)) != 1 { + err := errors.New("the authorization response state did not match the request; the login attempt may have been forged") + return callbackResult{err: err}, errorContent(err.Error()), http.StatusBadRequest + } + + if code := query.Get("error"); code != "" { + err := &AuthorizationError{Code: code, Description: query.Get("error_description")} + return callbackResult{err: err}, errorContent(err.Error()), http.StatusBadRequest + } + + code := query.Get("code") + if code == "" { + err := errors.New("the authorization response did not include an authorization code") + return callbackResult{err: err}, errorContent(err.Error()), http.StatusBadRequest + } + + return callbackResult{code: code}, successContent, http.StatusOK +} + +// generateCodeVerifier returns a PKCE code verifier: 32 random bytes encoded +// as 43 unreserved characters, within the 43-128 character range RFC 7636 +// section 4.1 allows. +func generateCodeVerifier() (string, error) { + return randomURLSafeString(32) +} + +func codeChallengeS256(verifier string) string { + sum := sha256.Sum256([]byte(verifier)) + return base64.RawURLEncoding.EncodeToString(sum[:]) +} + +func randomURLSafeString(byteLen int) (string, error) { + buf := make([]byte, byteLen) + if _, err := rand.Read(buf); err != nil { + return "", fmt.Errorf("generating random data: %w", err) + } + return base64.RawURLEncoding.EncodeToString(buf), nil +} diff --git a/internal/oauth/login_test.go b/internal/oauth/login_test.go new file mode 100644 index 000000000..2f7861685 --- /dev/null +++ b/internal/oauth/login_test.go @@ -0,0 +1,261 @@ +/* +Copyright 2018 The Doctl Authors All rights reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package oauth + +import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "net/url" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestLogin(t *testing.T) { + var challenge, verifier string + + tokenServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.NoError(t, r.ParseForm()) + verifier = r.FormValue("code_verifier") + assert.Equal(t, "auth-code", r.FormValue("code")) + + writeJSON(t, w, http.StatusOK, map[string]any{ + "access_token": "doo_v1_access", + "refresh_token": "dor_v1_refresh", + "expires_in": 3600, + }) + })) + defer tokenServer.Close() + + var authorizedRedirectURI string + // Stand in for the browser: read the authorization request, then call the + // redirect URI the way the authorization server would. + openURL := func(authURL string) error { + parsed, err := url.Parse(authURL) + require.NoError(t, err) + + query := parsed.Query() + assert.Equal(t, "client-id", query.Get("client_id")) + assert.Equal(t, responseTypeCode, query.Get("response_type")) + assert.Equal(t, CodeChallengeMethodS256, query.Get("code_challenge_method")) + assert.Equal(t, "read write", query.Get("scope")) + assert.NotEmpty(t, query.Get("state")) + assert.Empty(t, query.Get("resource"), "resource indicators are not requested") + + challenge = query.Get("code_challenge") + authorizedRedirectURI = query.Get("redirect_uri") + + return followRedirect(t, query.Get("redirect_uri"), url.Values{ + "code": {"auth-code"}, + "state": {query.Get("state")}, + }) + } + + token, err := Login(context.Background(), LoginOptions{ + Metadata: &ServerMetadata{AuthorizationEndpoint: "https://cloud.example.com/v1/oauth/authorize", TokenEndpoint: tokenServer.URL}, + ClientID: "client-id", + Scopes: []string{"read", "write"}, + HTTPClient: tokenServer.Client(), + openURL: openURL, + }) + require.NoError(t, err) + + assert.Equal(t, "doo_v1_access", token.AccessToken) + assert.Equal(t, "dor_v1_refresh", token.RefreshToken) + assert.Equal(t, challenge, codeChallengeS256(verifier), "the verifier must hash to the challenge sent to the authorization server") + assert.Regexp(t, `^http://127\.0\.0\.1:\d+/callback$`, authorizedRedirectURI) +} + +func TestLoginReportsAuthorizationErrors(t *testing.T) { + tests := []struct { + name string + params func(state string) url.Values + expectedError string + invalidClient bool + }{ + { + name: "access denied", + params: func(state string) url.Values { + return url.Values{ + "error": {"access_denied"}, + "error_description": {"the user declined the request"}, + "state": {state}, + } + }, + expectedError: "the user declined the request", + }, + { + name: "unknown client", + params: func(state string) url.Values { + return url.Values{"error": {"invalid_client"}, "state": {state}} + }, + expectedError: "invalid_client", + invalidClient: true, + }, + { + name: "mismatched state", + params: func(string) url.Values { + return url.Values{"code": {"auth-code"}, "state": {"not-the-state"}} + }, + expectedError: "state did not match", + }, + { + name: "missing code", + params: func(state string) url.Values { + return url.Values{"state": {state}} + }, + expectedError: "did not include an authorization code", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + openURL := func(authURL string) error { + parsed, err := url.Parse(authURL) + require.NoError(t, err) + + query := parsed.Query() + return followRedirect(t, query.Get("redirect_uri"), test.params(query.Get("state"))) + } + + _, err := Login(context.Background(), LoginOptions{ + Metadata: &ServerMetadata{AuthorizationEndpoint: "https://cloud.example.com/v1/oauth/authorize", TokenEndpoint: "https://cloud.example.com/v1/oauth/token"}, + ClientID: "client-id", + HTTPClient: failingClient(t), + openURL: openURL, + }) + + require.Error(t, err) + assert.Contains(t, err.Error(), test.expectedError) + if test.invalidClient { + assert.ErrorIs(t, err, ErrInvalidClient) + } else { + assert.NotErrorIs(t, err, ErrInvalidClient) + } + }) + } +} + +func TestCallbackPageEscapesServerSuppliedText(t *testing.T) { + const payload = `` + + handler := &callbackHandler{state: "the-state", results: make(chan callbackResult, 1)} + request := httptest.NewRequest(http.MethodGet, callbackPath+"?"+url.Values{ + "error": {payload}, + "error_description": {`">`}, + "state": {"the-state"}, + }.Encode(), nil) + + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, request) + + page := recorder.Body.String() + assert.NotContains(t, page, "