diff --git a/cloud/auth/auth.go b/cloud/auth/auth.go index b04503e88..870ecff83 100644 --- a/cloud/auth/auth.go +++ b/cloud/auth/auth.go @@ -423,27 +423,29 @@ func Login(domain, token string, astroV1Client astrov1.APIClient, out io.Writer, return nil } -// Logout logs a user out of the docker registry. Will need to logout of Astro next. -func Logout(domain string, out io.Writer) { - c, _ := context.GetContext(domain) - - err = c.SetContextKey("token", "") +// Logout logs a user out of an Astro domain. It clears every credential held in +// the domain's context (access token, refresh token, email and expiry) so no +// live token is left behind in the config file, and unsets the current context +// when that domain was the current one. +func Logout(domain string, out io.Writer) error { + c, err := context.GetContext(domain) if err != nil { - return + return fmt.Errorf("failed to find a login for %s: %w", domain, err) } - err = c.SetContextKey("user_email", "") - if err != nil { - return + + if err := c.ClearCredentials(); err != nil { + return fmt.Errorf("failed to clear the credentials for %s: %w", domain, err) } - // remove the current context - err = config.ResetCurrentContext() - if err != nil { - fmt.Fprintln(out, "Failed to reset current context: ", err.Error()) - return + // Logging out of another domain leaves the current one selected + if current, err := config.GetCurrentDomain(); err == nil && current == domain { + if err := config.ResetCurrentContext(); err != nil { + return fmt.Errorf("failed to reset the current context: %w", err) + } } fmt.Fprintln(out, "Successfully logged out of Astronomer") + return nil } func FetchDomainAuthConfig(domain string) (Config, error) { diff --git a/cloud/auth/auth_test.go b/cloud/auth/auth_test.go index 753a4704f..b58600147 100644 --- a/cloud/auth/auth_test.go +++ b/cloud/auth/auth_test.go @@ -948,40 +948,75 @@ func TestLogin(t *testing.T) { } func TestLogout(t *testing.T) { - testUtil.InitTestConfig(testUtil.LocalPlatform) - t.Run("success", func(t *testing.T) { + // loggedIn gives the context every credential a login writes + loggedIn := func(t *testing.T, c config.Context) { + t.Helper() + assert.NoError(t, c.SetContextKey("user_email", "test.user@astronomer.io")) + assert.NoError(t, c.SetContextKey("token", "Bearer some-token")) + assert.NoError(t, c.SetContextKey("refreshtoken", "some-refresh-token")) + assert.NoError(t, c.SetExpiresIn(3600)) + } + stored := func(t *testing.T, domain string) config.Context { + t.Helper() + c, err := (&config.Context{Domain: domain}).GetContext() + assert.NoError(t, err) + return c + } + + t.Run("clears every credential and the current context", func(t *testing.T) { + testUtil.InitTestConfig(testUtil.LocalPlatform) + c, err := config.GetCurrentContext() + assert.NoError(t, err) + loggedIn(t, c) + buf := new(bytes.Buffer) - Logout("astronomer.io", buf) + assert.NoError(t, Logout(c.Domain, buf)) + + after := stored(t, c.Domain) + assert.Empty(t, after.Token) + assert.Empty(t, after.RefreshToken) + assert.Empty(t, after.UserEmail) + expiresIn, err := after.GetExpiresIn() + assert.NoError(t, err) + assert.True(t, expiresIn.IsZero(), "expiry left behind: %v", expiresIn) + _, err = config.GetCurrentDomain() + assert.ErrorIs(t, err, config.ErrGetHomeString) assert.Equal(t, "Successfully logged out of Astronomer\n", buf.String()) }) - t.Run("success_with_email", func(t *testing.T) { - assertions := func(expUserEmail string, expToken string) { - contexts, err := config.GetContexts() - assert.NoError(t, err) - context := contexts.Contexts["localhost"] - - assert.NoError(t, err) - assert.Equal(t, expUserEmail, context.UserEmail) - assert.Equal(t, expToken, context.Token) - } + t.Run("logging out of another domain keeps the current one", func(t *testing.T) { testUtil.InitTestConfig(testUtil.LocalPlatform) - c, err := config.GetCurrentContext() - assert.NoError(t, err) - err = c.SetContextKey("user_email", "test.user@astronomer.io") + current, err := config.GetCurrentContext() assert.NoError(t, err) - err = c.SetContextKey("token", "Bearer some-token") + loggedIn(t, current) + other := config.Context{Domain: "astronomer-dev.io", Token: "Bearer other-token", RefreshToken: "other-refresh-token"} + assert.NoError(t, other.SetContext()) + + assert.NoError(t, Logout(other.Domain, new(bytes.Buffer))) + + assert.Empty(t, stored(t, other.Domain).RefreshToken) + domain, err := config.GetCurrentDomain() assert.NoError(t, err) - // test before - assertions("test.user@astronomer.io", "Bearer some-token") + assert.Equal(t, current.Domain, domain) + assert.Equal(t, "some-refresh-token", stored(t, current.Domain).RefreshToken) + }) - // log out - c, err = config.GetCurrentContext() + t.Run("an unknown domain fails without reporting success", func(t *testing.T) { + testUtil.InitTestConfig(testUtil.LocalPlatform) + current, err := config.GetCurrentContext() assert.NoError(t, err) - Logout(c.Domain, os.Stdout) + loggedIn(t, current) - // test after logout - assertions("", "") + buf := new(bytes.Buffer) + err = Logout("never-logged-in.io", buf) + assert.ErrorContains(t, err, "never-logged-in.io") + assert.Empty(t, buf.String()) + // no stub context is written, and the real login is untouched + assert.False(t, (&config.Context{Domain: "never-logged-in.io"}).ContextExists()) + assert.Equal(t, "some-refresh-token", stored(t, current.Domain).RefreshToken) + domain, err := config.GetCurrentDomain() + assert.NoError(t, err) + assert.Equal(t, current.Domain, domain) }) } diff --git a/cmd/auth.go b/cmd/auth.go index 4a152dd0d..c7267630c 100644 --- a/cmd/auth.go +++ b/cmd/auth.go @@ -83,10 +83,9 @@ func logout(cmd *cobra.Command, args []string, out io.Writer) error { cmd.SilenceUsage = true if context.IsCloudDomain(domain) { - cloudLogout(domain, out) - } else { - apcLogout(domain) + return cloudLogout(domain, out) } + apcLogout(domain) return nil } diff --git a/cmd/auth_test.go b/cmd/auth_test.go index 13d422ce9..b907719fd 100644 --- a/cmd/auth_test.go +++ b/cmd/auth_test.go @@ -2,6 +2,7 @@ package cmd import ( "bytes" + "errors" "io" "os" @@ -64,8 +65,13 @@ func (s *CmdSuite) TestLogout() { localDomain := "localhost" apcDomain := "astronomer_dev.com" - cloudLogout = func(domain string, out io.Writer) { + origCloudLogout, origAPCLogout := cloudLogout, apcLogout + defer func() { cloudLogout, apcLogout = origCloudLogout, origAPCLogout }() + + var cloudLogoutErr error + cloudLogout = func(domain string, out io.Writer) error { s.Equal(localDomain, domain) + return cloudLogoutErr } apcLogout = func(domain string) { s.Equal(apcDomain, domain) @@ -75,6 +81,12 @@ func (s *CmdSuite) TestLogout() { err := logout(&cobra.Command{}, []string{localDomain}, os.Stdout) s.NoError(err) + // a cloud logout that fails fails the command, so it exits non-zero + cloudLogoutErr = errors.New("failed to clear the credentials") + err = logout(&cobra.Command{}, []string{localDomain}, os.Stdout) + s.ErrorIs(err, cloudLogoutErr) + cloudLogoutErr = nil + // software logout success err = logout(&cobra.Command{}, []string{apcDomain}, os.Stdout) s.NoError(err) diff --git a/config/context.go b/config/context.go index 258c91d05..cce5b699f 100644 --- a/config/context.go +++ b/config/context.go @@ -159,16 +159,40 @@ func (c *Context) SetContextKey(key, value string) error { // a partial struct with every unset field zeroed out. // See https://github.com/spf13/viper/issues/1106. func setContextField(cKey, field string, value interface{}) error { + return updateContextMap(cKey, func(ctxMap map[string]interface{}) { ctxMap[field] = value }) +} + +// updateContextMap applies update to the context's map and persists the config +// in one write, for the reason setContextField gives. +func updateContextMap(cKey string, update func(map[string]interface{})) error { parentPath := fmt.Sprintf("%s.%s", contextsKey, cKey) ctxMap := viperHome.GetStringMap(parentPath) if ctxMap == nil { ctxMap = map[string]interface{}{} } - ctxMap[field] = value + update(ctxMap) viperHome.Set(parentPath, ctxMap) return saveConfig(viperHome, HomeConfigFile) } +// ClearCredentials empties every credential the context holds: the access +// token, the refresh token, the user's email and the token's expiry. It is one +// write, so a logout that fails to save leaves the context as it was rather +// than with the long-lived refresh token still in place. +func (c *Context) ClearCredentials() error { + cKey, err := c.GetContextKey() + if err != nil { + return err + } + return updateContextMap(cKey, func(ctxMap map[string]interface{}) { + ctxMap["token"] = "" + ctxMap["refreshtoken"] = "" + ctxMap["user_email"] = "" + // viper lowercases keys, so SetExpiresIn's "ExpiresIn" is stored as this + delete(ctxMap, "expiresin") + }) +} + // set organization id and short name in context config func (c *Context) SetOrganizationContext(orgID, orgProduct string) error { err := c.SetContextKey("organization", orgID) // c.Organization