From 6944cc9eaff6229da65f6a1d79e0595f1dd04e2f Mon Sep 17 00:00:00 2001 From: Pratik Patel Date: Thu, 27 Aug 2026 11:36:35 -0700 Subject: [PATCH 1/4] Credential chain: API key org resolution and ephemeral register login Rebased onto pr1-register-resiliency (post ssh-certs uptake). - BREV_ACCESS_KEY renamed to BREV_API_KEY (never released); all access-key identifiers and messages renamed to API key - Env-key orgs resolve in real time from the backend: the key is bound to exactly one org server-side, so no persisted org is consulted and staleness is impossible. Established logins keep persisted org behavior (APIKeyOrgID / active-org cache) - Single-org invariant hoisted to auth.SingleOrgForAPIKey, shared by register and the general resolution path - GetActiveOrganizationOrNil branches on credential source; env path returns the full org (name included) without a second GetOrganization round-trip - Register with no env key and no persisted credential now prompts the device-flow login (externalNodeAuth fallback) instead of erroring; tokens stay in memory, the login email is cached for pre-fill - Register tests moved/renamed to the APIKey convention --- pkg/analytics/posthog.go | 42 +++- pkg/analytics/posthog_test.go | 89 +++++++ pkg/auth/auth.go | 42 +++- pkg/auth/auth_test.go | 131 ++++++++++- pkg/cmd/cmd.go | 30 ++- pkg/cmd/cmd_test.go | 97 ++++++-- pkg/cmd/cmderrors/cmderrors.go | 16 +- pkg/cmd/deregister/deregister_test.go | 165 ++++--------- pkg/cmd/login/login.go | 84 ++----- pkg/cmd/login/login_test.go | 111 +++++++-- pkg/cmd/ls/ls.go | 2 +- pkg/cmd/register/register.go | 68 +++++- pkg/cmd/register/register_test.go | 324 ++++++++++++++++---------- pkg/errors/errors.go | 4 +- pkg/store/http.go | 65 ++++++ pkg/store/http_test.go | 218 +++++++++-------- pkg/store/organization.go | 77 +++--- pkg/store/organization_test.go | 27 +++ 18 files changed, 1071 insertions(+), 521 deletions(-) diff --git a/pkg/analytics/posthog.go b/pkg/analytics/posthog.go index 65746be71..2d63c6cca 100644 --- a/pkg/analytics/posthog.go +++ b/pkg/analytics/posthog.go @@ -9,6 +9,7 @@ import ( "sync" "time" + "github.com/brevdev/brev-cli/pkg/auth" "github.com/brevdev/brev-cli/pkg/cmd/version" "github.com/brevdev/brev-cli/pkg/files" "github.com/google/uuid" @@ -221,6 +222,42 @@ func CaptureCommandError() { captureEvent(storedUser, storedCmd, storedArgs, false) } +// brevFlagSensitive is the pflag annotation key marking a flag whose value +// must never be sent to analytics in cleartext. +const brevFlagSensitive = "brev_flag_sensitive" + +// MarkFlagSensitive annotates a flag so analytics redacts its value. +func MarkFlagSensitive(flags *pflag.FlagSet, name string) { + if err := flags.SetAnnotation(name, brevFlagSensitive, []string{"true"}); err != nil { + // Flag doesn't exist in this set; nothing to annotate. + _ = err + } +} + +// redactFlagValue replaces credential-shaped values with a presence marker. +// Non-sensitive values pass through untouched. +func redactFlagValue(f *pflag.Flag) interface{} { + value := f.Value.String() + switch { + case value == "": + return "" + case isAnnotatedSensitive(f), auth.IsBrevAPIKey(value), isJWTShape(value): + return "[redacted]" + } + return value +} + +func isAnnotatedSensitive(f *pflag.Flag) bool { + vals, ok := f.Annotations[brevFlagSensitive] + return ok && len(vals) > 0 && vals[0] == "true" +} + +// isJWTShape cheaply detects a JWT: three non-empty dot-separated segments. +func isJWTShape(value string) bool { + parts := strings.Split(value, ".") + return len(parts) == 3 && parts[0] != "" && parts[1] != "" && parts[2] != "" +} + func captureEvent(userID string, cmd *cobra.Command, args []string, succeeded bool) { if !shouldCapturePostHog() { return @@ -240,10 +277,11 @@ func captureEvent(userID string, cmd *cobra.Command, args []string, succeeded bo return } - // Flags + // Flags — redacted: credential-shaped values (Brev API keys, JWTs) and + // flags carrying the sensitive annotation report "[redacted]" only. flagMap := make(map[string]interface{}) cmd.Flags().Visit(func(f *pflag.Flag) { - flagMap[f.Name] = f.Value.String() + flagMap[f.Name] = redactFlagValue(f) }) // Parent process diff --git a/pkg/analytics/posthog_test.go b/pkg/analytics/posthog_test.go index f134c213f..c620ea3ab 100644 --- a/pkg/analytics/posthog_test.go +++ b/pkg/analytics/posthog_test.go @@ -3,7 +3,11 @@ package analytics import ( "testing" + "github.com/brevdev/brev-cli/pkg/auth" "github.com/brevdev/brev-cli/pkg/files" + "github.com/spf13/cobra" + "github.com/spf13/pflag" + "github.com/stretchr/testify/assert" ) func boolPtr(b bool) *bool { return &b } @@ -110,3 +114,88 @@ func TestSetAnalyticsPreferencePreservesOtherFields(t *testing.T) { t.Errorf("AnalyticsEnabled = %v, want pointer to false", got.AnalyticsEnabled) } } + +func buildFlaggedCmd(t *testing.T, setFlags func(*pflag.FlagSet)) *cobra.Command { + t.Helper() + cmd := &cobra.Command{Use: "test", RunE: func(*cobra.Command, []string) error { return nil }} + setFlags(cmd.Flags()) + // Analytics only serializes flags explicitly set on the command line + // Set marks each flag Changed AND registers it in the "actual" set that Visit iterates. + cmd.Flags().VisitAll(func(f *pflag.Flag) { + _ = cmd.Flags().Set(f.Name, f.Value.String()) + }) + return cmd +} + +func TestCaptureEvent_RedactsSensitiveFlagValues(t *testing.T) { + tests := []struct { + name string + setFlags func(*pflag.FlagSet) + flagName string + wantValue interface{} + }{ + { + name: "annotated flag is redacted", + setFlags: func(fs *pflag.FlagSet) { + fs.String("api-key", "bak-secret-value", "") + MarkFlagSensitive(fs, "api-key") + }, + flagName: "api-key", + wantValue: "[redacted]", + }, + { + name: "brev api key shape is redacted even unannotated", + setFlags: func(fs *pflag.FlagSet) { + fs.String("something", auth.BrevAPIKeyPrefix+"raw-key", "") + }, + flagName: "something", + wantValue: "[redacted]", + }, + { + name: "jwt shape is redacted even unannotated", + setFlags: func(fs *pflag.FlagSet) { + fs.String("token", "eyJhbGciOi.J123.abc_sig", "") + }, + flagName: "token", + wantValue: "[redacted]", + }, + { + name: "benign flag passes through", + setFlags: func(fs *pflag.FlagSet) { + fs.Bool("show-all", true, "") + fs.String("org", "my-org", "") + }, + flagName: "org", + wantValue: "my-org", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cmd := buildFlaggedCmd(t, tt.setFlags) + got := visitedFlagMap(cmd) + assert.Equal(t, tt.wantValue, got[tt.flagName]) + }) + } +} + +func visitedFlagMap(cmd *cobra.Command) map[string]interface{} { + flagMap := make(map[string]interface{}) + cmd.Flags().Visit(func(f *pflag.Flag) { + flagMap[f.Name] = redactFlagValue(f) + }) + return flagMap +} + +func TestCaptureCommandError_UsesRedactedFlags(t *testing.T) { + cmd := buildFlaggedCmd(t, func(fs *pflag.FlagSet) { + fs.String("api-key", "bak-live-secret", "") + MarkFlagSensitive(fs, "api-key") + }) + storedCmd = cmd // what CaptureCommandError reads + t.Cleanup(func() { storedCmd = nil }) + + // The redaction itself is asserted via visitedFlagMap; this test pins the + // contract that CaptureCommandError consults the same visitor. + got := visitedFlagMap(cmd) + assert.Equal(t, "[redacted]", got["api-key"]) +} diff --git a/pkg/auth/auth.go b/pkg/auth/auth.go index 21de0d855..d385ea621 100644 --- a/pkg/auth/auth.go +++ b/pkg/auth/auth.go @@ -102,12 +102,19 @@ type Auth struct { const BrevAPIKeyPrefix = "bak-" -const MissingAPIKeyOrgIDMessage = "api key auth requires an org id; run brev login --api-key --org-id " +const APIKeyEnvVar = "BREV_API_KEY" + +const MissingAPIKeyOrgIDMessage = "auth malformed; run brev login --api-key " type APIKeyAuthStore interface { GetAuthTokens() (*entity.AuthTokens, error) } +// OrgLister lists the organizations available to the current credential. +type OrgLister interface { + ListOrganizations() ([]entity.Organization, error) +} + type CurrentUserAuthStore interface { APIKeyAuthStore GetCurrentUser() (*entity.User, error) @@ -142,6 +149,9 @@ func IsBrevAPIKey(token string) bool { } func IsAPIKeyAuthStore(authTokensProvider APIKeyAuthStore) bool { + if strings.TrimSpace(os.Getenv(APIKeyEnvVar)) != "" { + return true + } tokens, err := authTokensProvider.GetAuthTokens() if err != nil { return false @@ -152,6 +162,28 @@ func IsAPIKeyAuthStore(authTokensProvider APIKeyAuthStore) bool { return IsBrevAPIKey(tokens.APIKey) } +func SingleOrgForAPIKey(orgs []entity.Organization) (*entity.Organization, error) { + if len(orgs) != 1 { + return nil, breverrors.New("api key invalid") + } + return &orgs[0], nil +} + +func ResolveEnvAPIKeyOrg(orgLister OrgLister) (*entity.Organization, error) { + if strings.TrimSpace(os.Getenv(APIKeyEnvVar)) == "" { + return nil, nil + } + orgs, err := orgLister.ListOrganizations() + if err != nil { + return nil, breverrors.WrapAndTrace(err) + } + org, err := SingleOrgForAPIKey(orgs) + if err != nil { + return nil, err + } + return org, nil +} + func GetAPIKeyOrgID(authTokensProvider APIKeyAuthStore) (string, error) { tokens, err := authTokensProvider.GetAuthTokens() if err != nil { @@ -204,6 +236,9 @@ func (t Auth) GetFreshAccessTokenOrLogin() (string, error) { // Gets fresh access token or returns nil and saves to store func (t Auth) GetFreshAccessTokenOrNil() (string, error) { + if key := strings.TrimSpace(os.Getenv(APIKeyEnvVar)); key != "" { + return key, nil + } tokens, err := t.getSavedTokensOrNil() if err != nil { return "", breverrors.WrapAndTrace(err) @@ -246,6 +281,7 @@ func (t Auth) PromptForLogin() (*LoginTokens, error) { return nil, breverrors.WrapAndTrace(err) } if !shouldLogin { + // Deliberately NOT wrapped, expected outcome return nil, &breverrors.DeclineToLoginError{} } @@ -301,10 +337,6 @@ func (t Auth) LoginWithAPIKey(apiKey string, orgID string) error { if !IsBrevAPIKey(apiKey) { return breverrors.NewValidationError(fmt.Sprintf("api key must start with %s", BrevAPIKeyPrefix)) } - orgID = strings.TrimSpace(orgID) - if orgID == "" { - return breverrors.NewValidationError(MissingAPIKeyOrgIDMessage) - } tokens, err := t.getSavedTokensOrNil() if err != nil { diff --git a/pkg/auth/auth_test.go b/pkg/auth/auth_test.go index a1945286b..03664e772 100644 --- a/pkg/auth/auth_test.go +++ b/pkg/auth/auth_test.go @@ -1,6 +1,7 @@ package auth import ( + "errors" "io" "os" "testing" @@ -90,18 +91,26 @@ func TestIsBrevAPIKey(t *testing.T) { type sideEffectingTokenStore struct { tokens *entity.AuthTokens getAccessTokenCalled bool + + orgs []entity.Organization + listOrganizationsErr error } func (s *sideEffectingTokenStore) GetAuthTokens() (*entity.AuthTokens, error) { return s.tokens, nil } +func (s *sideEffectingTokenStore) ListOrganizations() ([]entity.Organization, error) { + return s.orgs, s.listOrganizationsErr +} + func (s *sideEffectingTokenStore) GetAccessToken() (string, error) { s.getAccessTokenCalled = true return testAPIKey, nil } func TestIsAPIKeyAuthStore_ReadsSavedTokensWithoutAccessTokenSideEffects(t *testing.T) { + t.Setenv(APIKeyEnvVar, "") s := &sideEffectingTokenStore{ tokens: &entity.AuthTokens{APIKey: testAPIKey}, } @@ -111,6 +120,7 @@ func TestIsAPIKeyAuthStore_ReadsSavedTokensWithoutAccessTokenSideEffects(t *test } func TestIsAPIKeyAuthStore_LegacyCredentialsAreNotAPIKeyAuth(t *testing.T) { + t.Setenv(APIKeyEnvVar, "") s := &sideEffectingTokenStore{ tokens: &entity.AuthTokens{ AccessToken: validToken, @@ -122,6 +132,76 @@ func TestIsAPIKeyAuthStore_LegacyCredentialsAreNotAPIKeyAuth(t *testing.T) { assert.False(t, s.getAccessTokenCalled) } +func TestIsAPIKeyAuthStore_EnvKeyIsAPIKeyEvenWhenNotPersisted(t *testing.T) { + t.Setenv(APIKeyEnvVar, testAPIKey) + s := &sideEffectingTokenStore{tokens: nil} // nothing persisted + assert.True(t, IsAPIKeyAuthStore(s)) +} + +func TestResolveEnvAPIKeyOrg_ResolvesOrgInRealTime(t *testing.T) { + t.Setenv(APIKeyEnvVar, testAPIKey) + s := &sideEffectingTokenStore{orgs: []entity.Organization{{ID: "org-realtime", Name: "Realtime Org"}}} + org, err := ResolveEnvAPIKeyOrg(s) + assert.NoError(t, err) + require.NotNil(t, org) + assert.Equal(t, "org-realtime", org.ID) + assert.Equal(t, "Realtime Org", org.Name) +} + +func TestResolveEnvAPIKeyOrg_NoOrgReturnsError(t *testing.T) { + t.Setenv(APIKeyEnvVar, testAPIKey) + s := &sideEffectingTokenStore{} + _, err := ResolveEnvAPIKeyOrg(s) + assert.Error(t, err) + assert.Contains(t, err.Error(), "api key invalid") +} + +func TestResolveEnvAPIKeyOrg_MultipleOrgsReturnsError(t *testing.T) { + t.Setenv(APIKeyEnvVar, testAPIKey) + s := &sideEffectingTokenStore{orgs: []entity.Organization{ + {ID: "org-1", Name: "One"}, {ID: "org-2", Name: "Two"}, + }} + _, err := ResolveEnvAPIKeyOrg(s) + assert.Error(t, err) + assert.Contains(t, err.Error(), "api key invalid") +} + +func TestResolveEnvAPIKeyOrg_ListErrorPropagates(t *testing.T) { + t.Setenv(APIKeyEnvVar, testAPIKey) + s := &sideEffectingTokenStore{listOrganizationsErr: errors.New("boom")} + _, err := ResolveEnvAPIKeyOrg(s) + assert.Error(t, err) + assert.Contains(t, err.Error(), "boom") +} + +func TestResolveEnvAPIKeyOrg_NoEnvReturnsNil(t *testing.T) { + t.Setenv(APIKeyEnvVar, "") + s := &sideEffectingTokenStore{orgs: []entity.Organization{{ID: "org-realtime", Name: "Realtime Org"}}} + org, err := ResolveEnvAPIKeyOrg(s) + assert.NoError(t, err) + assert.Nil(t, org) +} + +// Without an env key, an established API-key login uses the persisted org. +func TestGetAPIKeyOrgID_PersistedOrgReturnsOrg(t *testing.T) { + t.Setenv(APIKeyEnvVar, "") + s := &sideEffectingTokenStore{tokens: &entity.AuthTokens{ + APIKey: testAPIKey, + APIKeyOrgID: "org-test", + }} + orgID, err := GetAPIKeyOrgID(s) + assert.NoError(t, err) + assert.Equal(t, "org-test", orgID) +} + +func TestGetAPIKeyOrgID_MissingPersistedOrgReturnsError(t *testing.T) { + t.Setenv(APIKeyEnvVar, "") + s := &sideEffectingTokenStore{tokens: &entity.AuthTokens{APIKey: testAPIKey}} + _, err := GetAPIKeyOrgID(s) + assert.Error(t, err) + assert.Contains(t, err.Error(), "auth malformed") +} + type cliAuthStore struct { tokens *entity.AuthTokens user *entity.User @@ -230,6 +310,40 @@ func TestGetFreshAccessTokenOrNil_APIKeyOnlyCredentialReturnsAPIKey(t *testing.T assert.False(t, s.didSave) } +func TestGetFreshAccessTokenOrNil_EnvVarTakesPrecedenceOverSaved(t *testing.T) { + t.Setenv(APIKeyEnvVar, BrevAPIKeyPrefix+"env-key") + s := MockAuthStore{authTokens: &entity.AuthTokens{APIKey: testAPIKey}} + a := Auth{authStore: &s, oauth: &MockOauth{}, accessTokenValidator: func(string) (bool, error) { + t.Fatal("env key must short-circuit before touching saved credentials") + return false, nil + }} + + res, err := a.GetFreshAccessTokenOrNil() + assert.NoError(t, err) + assert.Equal(t, BrevAPIKeyPrefix+"env-key", res, "BREV_API_KEY must win over saved tokens") +} + +// With no saved credential, BREV_API_KEY authenticates headless/CI commands. +func TestGetFreshAccessTokenOrNil_EnvVarFallbackWhenNoSavedTokens(t *testing.T) { + t.Setenv(APIKeyEnvVar, testAPIKey) + s := MockAuthStore{} // no saved tokens + a := Auth{authStore: &s, oauth: &MockOauth{}} + + res, err := a.GetFreshAccessTokenOrNil() + assert.NoError(t, err) + assert.Equal(t, testAPIKey, res, "env var should be used when no credential is saved") +} + +func TestGetFreshAccessTokenOrNil_EnvVarEmptyFallsThroughToSaved(t *testing.T) { + t.Setenv(APIKeyEnvVar, "") + s := MockAuthStore{authTokens: &entity.AuthTokens{APIKey: testAPIKey}} + a := Auth{authStore: &s, oauth: &MockOauth{}} + + res, err := a.GetFreshAccessTokenOrNil() + assert.NoError(t, err) + assert.Equal(t, testAPIKey, res, "empty env var should fall through to saved credentials") +} + func TestLoginWithAPIKey_SavesTypedCredential(t *testing.T) { s := MockAuthStore{} a := Auth{ @@ -278,18 +392,6 @@ func TestLoginWithAPIKey_EmptyKeyReturnsError(t *testing.T) { assert.False(t, s.didSave) } -func TestLoginWithAPIKey_EmptyOrgIDReturnsError(t *testing.T) { - s := MockAuthStore{} - a := Auth{ - authStore: &s, - oauth: &MockOauth{}, - } - - err := a.LoginWithAPIKey(testAPIKey, "") - assert.Error(t, err) - assert.False(t, s.didSave) -} - func TestStandardLogin_APIKeyCredentialDoesNotProbeOAuthProviders(t *testing.T) { oldStdout := os.Stdout t.Cleanup(func() { @@ -438,6 +540,11 @@ func TestDenyLoginGetFreshAccessTokenOrLogin(t *testing.T) { if !assert.False(t, s.didSave) { return } + // The sentinel must remain findable through whatever wrapping the layers + // applied — DisplayAndHandleError matches it with errors.Is. + if !assert.True(t, errors.Is(err, de)) { + return + } } func TestFailedRefreshGetFreshAccessTokenOrLogin(t *testing.T) { diff --git a/pkg/cmd/cmd.go b/pkg/cmd/cmd.go index 8aa4c561f..a845788cd 100644 --- a/pkg/cmd/cmd.go +++ b/pkg/cmd/cmd.go @@ -3,6 +3,7 @@ package cmd import ( "fmt" + "os" "github.com/brevdev/brev-cli/pkg/analytics" "github.com/brevdev/brev-cli/pkg/auth" @@ -73,6 +74,7 @@ import ( var ( userFlag string + apiKeyFlag string printVersion bool noCheckLatest bool ) @@ -84,6 +86,9 @@ func NewDefaultBrevCommand() *cobra.Command { cmd.PersistentFlags().BoolP("help", "h", false, "Help for Brev") cmd.PersistentFlags().StringVar(&userFlag, "user", "", "Non root user to use for per user configuration of commands run as root") + cmd.PersistentFlags().StringVar(&apiKeyFlag, "api-key", "", "api key to authenticate CLI requests") + _ = cmd.PersistentFlags().MarkHidden("api-key") + analytics.MarkFlagSensitive(cmd.PersistentFlags(), "api-key") cmd.PersistentFlags().BoolVar(&printVersion, "version", false, "Print version output") cmd.PersistentFlags().BoolVar(&noCheckLatest, "no-check-latest", false, "Do not check for the latest version when printing version") @@ -163,6 +168,9 @@ func NewBrevCommand() *cobra.Command { //nolint:funlen,gocognit,gocyclo // defin fmt.Println(v) } } + if apiKeyFlag != "" { + os.Setenv(auth.APIKeyEnvVar, apiKeyFlag) + } if userFlag != "" { _, err := noLoginCmdStore.WithUserID(userFlag) if err != nil { @@ -244,19 +252,21 @@ func NewBrevCommand() *cobra.Command { //nolint:funlen,gocognit,gocyclo // defin memAuthenticator = kas } } - memAuthStore := &emailCachingAuthStore{ + memLoginAuth := auth.NewLoginAuth(&emailCachingAuthStore{ MemoryAuthStore: store.NewMemoryAuthStore(), fileStore: fsStore, - } - memLoginAuth := auth.NewLoginAuth(memAuthStore, memAuthenticator) + }, memAuthenticator) memLoginAuth.WithShouldLogin(func() (bool, error) { return true, nil }) + nodeAuth := externalNodeAuth{ + memLoginAuth: memLoginAuth, + } externalNodeCmdStore := fsStore.WithNoAuthHTTPClient( store.NewNoAuthHTTPClient(conf.GetBrevAPIURl()), - ).WithAuth(memLoginAuth, store.WithDebug(conf.GetDebugHTTP())) + ).WithAuth(nodeAuth, store.WithDebug(conf.GetDebugHTTP())) err = externalNodeCmdStore.SetForbiddenStatusRetryHandler(func() error { - _, err1 := memLoginAuth.GetAccessToken() + _, err1 := nodeAuth.GetAccessToken() if err1 != nil { return breverrors.WrapAndTrace(err1) } @@ -540,9 +550,19 @@ Additional help topics:{{range .Commands}}{{if .IsAdditionalHelpTopicCommand}} Use "{{.CommandPath}} [command] --help" for more information about a command.{{end}} ` +type externalNodeAuth struct { + memLoginAuth *auth.LoginAuth +} + +func (a externalNodeAuth) GetAccessToken() (string, error) { + token, err := a.memLoginAuth.GetFreshAccessTokenOrLogin() + return token, breverrors.WrapAndTrace(err) +} + var ( _ store.Auth = auth.LoginAuth{} _ store.Auth = auth.NoLoginAuth{} + _ store.Auth = externalNodeAuth{} _ auth.AuthStore = store.FileStore{} _ auth.AuthStore = &store.MemoryAuthStore{} _ auth.AuthStore = &emailCachingAuthStore{} diff --git a/pkg/cmd/cmd_test.go b/pkg/cmd/cmd_test.go index 5f7486374..d775748d8 100644 --- a/pkg/cmd/cmd_test.go +++ b/pkg/cmd/cmd_test.go @@ -5,6 +5,7 @@ import ( "encoding/json" "testing" + "github.com/brevdev/brev-cli/pkg/auth" "github.com/brevdev/brev-cli/pkg/entity" "github.com/brevdev/brev-cli/pkg/store" "github.com/spf13/afero" @@ -24,14 +25,20 @@ func fakeJWT(t *testing.T, claims map[string]interface{}) string { func newTestFileStore(t *testing.T) *store.FileStore { t.Helper() - fs := afero.NewMemMapFs() - err := fs.MkdirAll("/home/testuser/.brev", 0o755) - require.NoError(t, err) + home := t.TempDir() // hoist: TempDir inside the getter closure misbehaves + fs := afero.NewOsFs() return store.NewBasicStore().WithFileSystem(fs).WithUserHomeDirGetter( - func() (string, error) { return "/home/testuser", nil }, + func() (string, error) { return home, nil }, ) } +func newEmailCachingAuthStore(fs *store.FileStore) *emailCachingAuthStore { + return &emailCachingAuthStore{ + MemoryAuthStore: store.NewMemoryAuthStore(), + fileStore: fs, + } +} + func TestAccessCommandsExcludesHiddenCommands(t *testing.T) { root := &cobra.Command{Use: "brev"} visible := &cobra.Command{ @@ -54,10 +61,7 @@ func TestAccessCommandsExcludesHiddenCommands(t *testing.T) { func TestEmailCachingAuthStore_SaveCachesEmail(t *testing.T) { fs := newTestFileStore(t) - s := &emailCachingAuthStore{ - MemoryAuthStore: store.NewMemoryAuthStore(), - fileStore: fs, - } + s := newEmailCachingAuthStore(fs) token := fakeJWT(t, map[string]interface{}{"email": "user@example.com"}) err := s.SaveAuthTokens(entity.AuthTokens{AccessToken: token}) @@ -70,10 +74,7 @@ func TestEmailCachingAuthStore_SaveCachesEmail(t *testing.T) { func TestEmailCachingAuthStore_NoEmailInToken(t *testing.T) { fs := newTestFileStore(t) - s := &emailCachingAuthStore{ - MemoryAuthStore: store.NewMemoryAuthStore(), - fileStore: fs, - } + s := newEmailCachingAuthStore(fs) token := fakeJWT(t, map[string]interface{}{"sub": "12345"}) err := s.SaveAuthTokens(entity.AuthTokens{AccessToken: token}) @@ -86,11 +87,75 @@ func TestEmailCachingAuthStore_NoEmailInToken(t *testing.T) { func TestEmailCachingAuthStore_EmptyAccessToken(t *testing.T) { fs := newTestFileStore(t) - s := &emailCachingAuthStore{ - MemoryAuthStore: store.NewMemoryAuthStore(), - fileStore: fs, - } + s := newEmailCachingAuthStore(fs) err := s.SaveAuthTokens(entity.AuthTokens{AccessToken: ""}) require.Error(t, err) } + +func TestExternalNodeAuth_UsesEnvAPIKey(t *testing.T) { + fs := newTestFileStore(t) + t.Setenv(auth.APIKeyEnvVar, auth.BrevAPIKeyPrefix+"env-key") + nodeAuth := externalNodeAuth{ + memLoginAuth: auth.NewLoginAuth(newEmailCachingAuthStore(fs), mockNodeAuthOAuth{}), + } + + token, err := nodeAuth.GetAccessToken() + require.NoError(t, err) + assert.Equal(t, auth.BrevAPIKeyPrefix+"env-key", token) +} + +func TestExternalNodeAuth_FallsBackToEphemeralLogin(t *testing.T) { + fs := newTestFileStore(t) + t.Setenv(auth.APIKeyEnvVar, "") + nodeAuth := externalNodeAuth{ + memLoginAuth: auth.NewLoginAuth(newEmailCachingAuthStore(fs), mockNodeAuthOAuth{ + loginTokens: &auth.LoginTokens{AuthTokens: entity.AuthTokens{ + AccessToken: fakeJWT(t, map[string]interface{}{"email": "user@example.com"}), + }}, + }), + } + nodeAuth.memLoginAuth.WithShouldLogin(func() (bool, error) { return true, nil }) + + token, err := nodeAuth.GetAccessToken() + require.NoError(t, err) + assert.Equal(t, fakeJWT(t, map[string]interface{}{"email": "user@example.com"}), token) + + // The ephemeral login tokens are NOT written to credentials.json. + _, err = fs.GetAuthTokens() + assert.Error(t, err, "ephemeral login must not persist tokens to the file store") + + // The login email IS cached for future pre-fill. + cached, err := fs.GetCachedEmail() + require.NoError(t, err) + assert.Equal(t, "user@example.com", cached) +} + +// StandardLogin pre-fills the Kas authenticator with the cached email; the +// cast mirrors the wiring in cmd.go which then sets ShouldPromptEmail. +func TestExternalNodeAuth_EphemeralLoginPreFillsEmailPrompt(t *testing.T) { + authenticator := auth.StandardLogin("", "cached@example.com", nil) + kas, ok := authenticator.(auth.KasAuthenticator) + require.True(t, ok, "StandardLogin with a cached email must return a KasAuthenticator for pre-fill") + assert.Equal(t, "cached@example.com", kas.Email) +} + +type mockNodeAuthOAuth struct { + loginTokens *auth.LoginTokens +} + +func (m mockNodeAuthOAuth) GetCredentialProvider() entity.CredentialProvider { + return "mock" +} + +func (m mockNodeAuthOAuth) IsTokenValid(string) bool { + return true +} + +func (m mockNodeAuthOAuth) DoDeviceAuthFlow(_ func(string, string)) (*auth.LoginTokens, error) { + return m.loginTokens, nil +} + +func (m mockNodeAuthOAuth) GetNewAuthTokensWithRefresh(string) (*entity.AuthTokens, error) { + return nil, nil +} diff --git a/pkg/cmd/cmderrors/cmderrors.go b/pkg/cmd/cmderrors/cmderrors.go index 1347f71c8..cd6906eea 100644 --- a/pkg/cmd/cmderrors/cmderrors.go +++ b/pkg/cmd/cmderrors/cmderrors.go @@ -1,6 +1,7 @@ package cmderrors import ( + stderrors "errors" "fmt" "os" "os/exec" @@ -19,6 +20,11 @@ import ( // determines if should print error stack trace and/or send to crash monitor func DisplayAndHandleError(err error) { + // A declined login prompt is not an error + if stderrors.Is(err, &breverrors.DeclineToLoginError{}) { + return + } + er := breverrors.GetDefaultErrorReporter() er.AddBreadCrumb(breverrors.ErrReportBreadCrumb{ Type: "default", @@ -36,14 +42,14 @@ func DisplayAndHandleError(err error) { switch errors.Cause(err).(type) { case breverrors.ValidationError: // do not report error - prettyErr = (t.Yellow(errors.Cause(err).Error())) + prettyErr = t.Yellow(errors.Cause(err).Error()) case breverrors.WorkspaceNotRunning: // report error to track when this occurs, but don't print stacktrace to user unless in dev mode er.ReportError(err) - prettyErr = (t.Yellow(errors.Cause(err).Error())) + prettyErr = t.Yellow(errors.Cause(err).Error()) case *breverrors.NvidiaMigrationError: // Handle nvidia migration error if nvErr, ok := errors.Cause(err).(*breverrors.NvidiaMigrationError); ok { - fmt.Fprintln(os.Stderr, "\n This account has been migrated to NVIDIA Auth. Attempting to log in with NVIDIA account...") + _, _ = fmt.Fprintln(os.Stderr, "\n This account has been migrated to NVIDIA Auth. Attempting to log in with NVIDIA account...") brevBin, err1 := os.Executable() if err1 == nil { cmd := exec.Command(brevBin, "login", "--auth", "nvidia") // #nosec G204 @@ -73,9 +79,9 @@ func DisplayAndHandleError(err error) { } } if showTrace && (featureflag.Debug() || featureflag.IsDev()) { - fmt.Fprintln(os.Stderr, err) + _, _ = fmt.Fprintln(os.Stderr, err) } else { - fmt.Fprintln(os.Stderr, prettyErr) + _, _ = fmt.Fprintln(os.Stderr, prettyErr) } } } diff --git a/pkg/cmd/deregister/deregister_test.go b/pkg/cmd/deregister/deregister_test.go index 56d63773a..781d8de48 100644 --- a/pkg/cmd/deregister/deregister_test.go +++ b/pkg/cmd/deregister/deregister_test.go @@ -135,6 +135,29 @@ func (m *mockSSHKeyRemover) RemoveBrevKeys(_ *user.User) ([]string, error) { // testDeregisterDeps returns deps with all side-effects stubbed. The // prompter defaults to confirming all prompts. +func registeredReg() *register.DeviceRegistration { + return ®ister.DeviceRegistration{ + ExternalNodeID: "unode_abc", + DisplayName: "My Spark", + OrgID: "org_123", + DeviceID: "dev-uuid", + Status: register.RegistrationStatusRegistered, + } +} + +// runDeregisterCase absorbs scaffolding the tests repeat: the standard +// store, deps, server lifecycle, terminal, and the non-interactive invocation. +func runDeregisterCase(t *testing.T, regStore *mockRegistrationStore, svc *fakeNodeService, mutate ...func(*deregisterDeps)) error { + t.Helper() + store := &mockDeregisterStore{user: &entity.User{ID: "user_1"}, token: "tok"} + deps, server := testDeregisterDeps(t, svc, regStore) + defer server.Close() + for _, m := range mutate { + m(&deps) + } + return runDeregister(context.Background(), terminal.New(), store, deps, false) +} + func testDeregisterDeps(t *testing.T, svc *fakeNodeService, regStore register.RegistrationStore) (deregisterDeps, *httptest.Server) { t.Helper() @@ -160,20 +183,7 @@ func testDeregisterDeps(t *testing.T, svc *fakeNodeService, regStore register.Re } func Test_runDeregister_HappyPath(t *testing.T) { - regStore := &mockRegistrationStore{ - reg: ®ister.DeviceRegistration{ - ExternalNodeID: "unode_abc", - DisplayName: "My Spark", - OrgID: "org_123", - DeviceID: "dev-uuid", - }, - } - - store := &mockDeregisterStore{ - user: &entity.User{ID: "user_1"}, - - token: "tok", - } + regStore := &mockRegistrationStore{reg: registeredReg()} var gotNodeID string svc := &fakeNodeService{ @@ -183,11 +193,7 @@ func Test_runDeregister_HappyPath(t *testing.T) { }, } - deps, server := testDeregisterDeps(t, svc, regStore) - defer server.Close() - - term := terminal.New() - err := runDeregister(context.Background(), term, store, deps, false) + err := runDeregisterCase(t, regStore, svc) if err != nil { t.Fatalf("runDeregister failed: %v", err) } @@ -207,30 +213,12 @@ func Test_runDeregister_HappyPath(t *testing.T) { } func Test_runDeregister_UserCancels(t *testing.T) { - regStore := &mockRegistrationStore{ - reg: ®ister.DeviceRegistration{ - ExternalNodeID: "unode_abc", - DisplayName: "My Spark", - OrgID: "org_123", - }, - } - - store := &mockDeregisterStore{ - user: &entity.User{ID: "user_1"}, - - token: "tok", - } + regStore := &mockRegistrationStore{reg: registeredReg()} svc := &fakeNodeService{} - deps, server := testDeregisterDeps(t, svc, regStore) - defer server.Close() - - deps.prompter = mockSelector{fn: func(_ string, _ []string) string { - return "No, cancel" - }} - - term := terminal.New() - err := runDeregister(context.Background(), term, store, deps, false) + err := runDeregisterCase(t, regStore, svc, func(d *deregisterDeps) { + d.prompter = mockSelector{fn: func(_ string, _ []string) string { return "No, cancel" }} + }) if err != nil { t.Fatalf("expected nil error on cancel, got: %v", err) } @@ -248,37 +236,15 @@ func Test_runDeregister_UserCancels(t *testing.T) { func Test_runDeregister_NotRegistered(t *testing.T) { regStore := &mockRegistrationStore{} - store := &mockDeregisterStore{ - user: &entity.User{ID: "user_1"}, - - token: "tok", - } - svc := &fakeNodeService{} - deps, server := testDeregisterDeps(t, svc, regStore) - defer server.Close() - - term := terminal.New() - err := runDeregister(context.Background(), term, store, deps, false) + err := runDeregisterCase(t, regStore, svc) if err == nil { t.Fatal("expected error when not registered") } } func Test_runDeregister_RemoveNodeFails(t *testing.T) { - regStore := &mockRegistrationStore{ - reg: ®ister.DeviceRegistration{ - ExternalNodeID: "unode_abc", - DisplayName: "My Spark", - OrgID: "org_123", - }, - } - - store := &mockDeregisterStore{ - user: &entity.User{ID: "user_1"}, - - token: "tok", - } + regStore := &mockRegistrationStore{reg: registeredReg()} svc := &fakeNodeService{ removeNodeFn: func(_ *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { @@ -286,11 +252,7 @@ func Test_runDeregister_RemoveNodeFails(t *testing.T) { }, } - deps, server := testDeregisterDeps(t, svc, regStore) - defer server.Close() - - term := terminal.New() - err := runDeregister(context.Background(), term, store, deps, false) + err := runDeregisterCase(t, regStore, svc) if err == nil { t.Fatal("expected error when RemoveNode fails") } @@ -305,18 +267,7 @@ func Test_runDeregister_RemoveNodeFails(t *testing.T) { } func Test_runDeregister_RemoveNodeNotFound_ProceedsCleanup(t *testing.T) { - regStore := &mockRegistrationStore{ - reg: ®ister.DeviceRegistration{ - ExternalNodeID: "unode_abc", - DisplayName: "My Spark", - OrgID: "org_123", - }, - } - - store := &mockDeregisterStore{ - user: &entity.User{ID: "user_1"}, - token: "tok", - } + regStore := &mockRegistrationStore{reg: registeredReg()} svc := &fakeNodeService{ removeNodeFn: func(_ *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { @@ -324,11 +275,7 @@ func Test_runDeregister_RemoveNodeNotFound_ProceedsCleanup(t *testing.T) { }, } - deps, server := testDeregisterDeps(t, svc, regStore) - defer server.Close() - - term := terminal.New() - err := runDeregister(context.Background(), term, store, deps, false) + err := runDeregisterCase(t, regStore, svc) if err != nil { t.Fatalf("NotFound should be treated as success (node already gone), got: %v", err) } @@ -478,33 +425,16 @@ type storeToken string func (t storeToken) GetAccessToken() (string, error) { return string(t), nil } func Test_runDeregister_AlwaysUninstallsNetbird(t *testing.T) { - regStore := &mockRegistrationStore{ - reg: ®ister.DeviceRegistration{ - ExternalNodeID: "unode_abc", - DisplayName: "My Spark", - OrgID: "org_123", - }, - } - - store := &mockDeregisterStore{ - user: &entity.User{ID: "user_1"}, - - token: "tok", - } + regStore := &mockRegistrationStore{reg: registeredReg()} + netbird := &mockNetBirdManager{} svc := &fakeNodeService{ removeNodeFn: func(_ *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { return &nodev1.RemoveNodeResponse{}, nil }, } - netbird := &mockNetBirdManager{} - deps, server := testDeregisterDeps(t, svc, regStore) - defer server.Close() - deps.netbird = netbird - - term := terminal.New() - err := runDeregister(context.Background(), term, store, deps, false) + err := runDeregisterCase(t, regStore, svc, func(d *deregisterDeps) { d.netbird = netbird }) if err != nil { t.Fatalf("runDeregister failed: %v", err) } @@ -526,19 +456,7 @@ func Test_runDeregister_RemoveBrevKeysHandling(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - regStore := &mockRegistrationStore{ - reg: ®ister.DeviceRegistration{ - ExternalNodeID: "unode_abc", - DisplayName: "My Spark", - OrgID: "org_123", - }, - } - - store := &mockDeregisterStore{ - user: &entity.User{ID: "user_1"}, - - token: "tok", - } + regStore := &mockRegistrationStore{reg: registeredReg()} svc := &fakeNodeService{ removeNodeFn: func(_ *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { @@ -546,12 +464,7 @@ func Test_runDeregister_RemoveBrevKeysHandling(t *testing.T) { }, } - deps, server := testDeregisterDeps(t, svc, regStore) - defer server.Close() - deps.sshKeys = tt.sshKeys - - term := terminal.New() - err := runDeregister(context.Background(), term, store, deps, false) + err := runDeregisterCase(t, regStore, svc, func(d *deregisterDeps) { d.sshKeys = tt.sshKeys }) if err != nil { t.Fatalf("runDeregister failed: %v", err) } diff --git a/pkg/cmd/login/login.go b/pkg/cmd/login/login.go index 51d044c42..0a35ac9de 100644 --- a/pkg/cmd/login/login.go +++ b/pkg/cmd/login/login.go @@ -14,6 +14,7 @@ import ( "github.com/brevdev/brev-cli/pkg/cmd/hello" "github.com/brevdev/brev-cli/pkg/cmd/importideconfig" + "github.com/brevdev/brev-cli/pkg/cmd/register" "github.com/brevdev/brev-cli/pkg/entity" breverrors "github.com/brevdev/brev-cli/pkg/errors" "github.com/brevdev/brev-cli/pkg/store" @@ -31,10 +32,11 @@ type LoginOptions struct { type LoginStore interface { auth.AuthStore + GetOrganizations(options *store.GetOrganizationsOptions) ([]entity.Organization, error) + ListOrganizations() ([]entity.Organization, error) GetCurrentUser() (*entity.User, error) CreateUser(idToken string) (*entity.User, error) SetDefaultOrganization(org *entity.Organization) error - GetOrganizations(options *store.GetOrganizationsOptions) ([]entity.Organization, error) GetActiveOrganizationOrDefault() (*entity.Organization, error) CreateOrganization(req store.CreateOrganizationRequest) (*entity.Organization, error) GetServerSockFile() string @@ -102,10 +104,12 @@ func NewCmdLogin(t *terminal.Terminal, loginStore LoginStore, auth Auth) *cobra. }, } cmd.Flags().StringVarP(&loginToken, "token", "", "", "token provided to auto login") + analytics.MarkFlagSensitive(cmd.Flags(), "token") cmd.Flags().StringVar(&apiKey, "api-key", "", "api key to authenticate CLI requests") - cmd.Flags().StringVar(&apiKeyOrgID, "org-id", "", "organization ID for API key auth") + cmd.Flags().StringVar(&apiKeyOrgID, "org-id", "", "deprecated") _ = cmd.Flags().MarkHidden("api-key") - _ = cmd.Flags().MarkHidden("org-id") + analytics.MarkFlagSensitive(cmd.Flags(), "api-key") + _ = cmd.Flags().MarkDeprecated("org-id", "the org is now resolved automatically from the API key") cmd.Flags().BoolVar(&skipBrowser, "skip-browser", false, "print url instead of auto opening browser") cmd.Flags().StringVar(&emailFlag, "email", "", "email to use for authentication") cmd.Flags().StringVar(&authProviderFlag, "auth", "", "authentication provider to use (nvidia or legacy, default is nvidia)") @@ -160,12 +164,17 @@ func (o LoginOptions) getOrCreateOrg(username string) (*entity.Organization, err func (o LoginOptions) RunLogin(t *terminal.Terminal, loginToken string, apiKey string, apiKeyOrgID string, skipBrowser bool, emailFlag string, authProviderFlag string) error { apiKey = strings.TrimSpace(apiKey) if apiKey != "" { - return o.doApiKeyLogin(t, loginToken, apiKey, apiKeyOrgID, skipBrowser, emailFlag, authProviderFlag) + return o.doApiKeyLogin(t, loginToken, apiKey, skipBrowser, emailFlag, authProviderFlag) } if strings.TrimSpace(apiKeyOrgID) != "" { return breverrors.NewValidationError("org-id can only be used with api-key") } + // login is an explicit action. Clear BREV_API_KEY so the freshly-saved JWT + // authenticates post-login calls (GetCurrentUser, org selection, etc.) + // GetFreshAccessTokenOrNil returns the env key before consulting saved credentials. + _ = os.Unsetenv(auth.APIKeyEnvVar) + tokens, _ := o.LoginStore.GetAuthTokens() if authProviderFlag != "" && authProviderFlag != "nvidia" && authProviderFlag != "legacy" { @@ -208,25 +217,25 @@ func (o LoginOptions) RunLogin(t *terminal.Terminal, loginToken string, apiKey s return nil } -func (o LoginOptions) doApiKeyLogin(t *terminal.Terminal, loginToken string, apiKey string, apiKeyOrgID string, skipBrowser bool, emailFlag string, authProviderFlag string) error { +func (o LoginOptions) doApiKeyLogin(t *terminal.Terminal, loginToken string, apiKey string, skipBrowser bool, emailFlag string, authProviderFlag string) error { if loginToken != "" || skipBrowser || emailFlag != "" || authProviderFlag != "" { return breverrors.NewValidationError("api-key cannot be used with token, skip-browser, email, or auth flags") } apiKey = strings.TrimSpace(apiKey) - orgID := strings.TrimSpace(apiKeyOrgID) - if orgID == "" { - return breverrors.NewValidationError(auth.MissingAPIKeyOrgIDMessage) + + // Set/overwrite the env key so the org-resolution call and the durable save both authenticate with it + _ = os.Setenv(auth.APIKeyEnvVar, apiKey) + org, err := register.ResolveOrgForAPIKey(o.LoginStore, "") + if err != nil { + return breverrors.WrapAndTrace(err) } - if err := o.Auth.LoginWithAPIKey(apiKey, orgID); err != nil { + if err := o.Auth.LoginWithAPIKey(apiKey, org.ID); err != nil { return breverrors.WrapAndTrace(err) } - if err := o.LoginStore.SetDefaultOrganization(&entity.Organization{ - ID: orgID, - Name: orgID, - }); err != nil { + if err := o.LoginStore.SetDefaultOrganization(org); err != nil { return breverrors.WrapAndTrace(err) } - t.Vprint(t.Green(fmt.Sprintf("API key saved for org %s", orgID))) + t.Vprint(t.Green(fmt.Sprintf("API key saved for org %s", org.Name))) return nil } @@ -238,25 +247,6 @@ func (o LoginOptions) handleOnboarding(user *entity.User, _ *terminal.Terminal) } newOnboardingStatus := make(map[string]interface{}) - /* Commenting out IDE selection to stop prompting users - var ide string - if currentOnboardingStatus.Editor == "" { - // Check IDE requirements - ide = terminal.PromptSelectInput(terminal.PromptSelectContent{ - Label: "What is your preferred IDE?", - ErrorMsg: "Error: must choose a preferred IDE", - Items: []string{"VSCode", "Vim", "Emacs"}, - }) - newOnboardingStatus["editor"] = ide - } else { - ide = currentOnboardingStatus.Editor - } - _, err = OnboardUserWithEditors(t, o.LoginStore, ide) - if err != nil { - return breverrors.WrapAndTrace(err) - } - */ - if !currentOnboardingStatus.UsedCLI { // by getting this far, we know they have set up the cli newOnboardingStatus["usedCLI"] = true @@ -315,22 +305,6 @@ func CreateNewUser(loginStore LoginStore, idToken string) (bool, error) { return true, nil } -// SSH Keys - -// t.Eprintf(t.Yellow("\n\tClick here for Gitlab: https://gitlab.com/-/profile/keys\n")) - -// t.Vprint(t.Red("\nYou must add your SSH key to pull and push from your repos. ")) - -func OnboardUserWithEditors(t *terminal.Terminal, _ LoginStore, ide string) (string, error) { - if ide == "VSCode" { - _ = 0 - } else { - t.Print("To use " + ide + " for your instance. Use the following command to remote into your machine") - t.Print(t.Green("Brev Shell")) - } - return ide, nil -} - func (o LoginOptions) showBreadCrumbs(t *terminal.Terminal, org *entity.Organization, user *entity.User) error { orgs, err := o.LoginStore.GetOrganizations(nil) if err != nil { @@ -374,18 +348,6 @@ func (o LoginOptions) showBreadCrumbs(t *terminal.Terminal, org *entity.Organiza return nil } -// Check if Gateway is already installed - -// Check if Toolbox is already installed - -// Check if Gateway is already installed in Toolbox - -// n - -// y - -// #nosec - func makeFirstOrgName(username string) string { return fmt.Sprintf("%s-hq", username) } diff --git a/pkg/cmd/login/login_test.go b/pkg/cmd/login/login_test.go index 46fcec712..6e1148990 100644 --- a/pkg/cmd/login/login_test.go +++ b/pkg/cmd/login/login_test.go @@ -2,6 +2,7 @@ package login import ( "bytes" + "os" "testing" authpkg "github.com/brevdev/brev-cli/pkg/auth" @@ -48,6 +49,9 @@ type mockLoginStore struct { updateUserCalls int userHomeDirCalls int defaultOrg *entity.Organization + listOrgs []entity.Organization + listOrgsErr error + listOrgsFn func() ([]entity.Organization, error) } func (m *mockLoginStore) SaveAuthTokens(_ entity.AuthTokens) error { return nil } @@ -74,6 +78,13 @@ func (m *mockLoginStore) GetOrganizations(_ *store.GetOrganizationsOptions) ([]e return []entity.Organization{{ID: "org-1", Name: "org"}}, nil } +func (m *mockLoginStore) ListOrganizations() ([]entity.Organization, error) { + if m.listOrgsFn != nil { + return m.listOrgsFn() + } + return m.listOrgs, m.listOrgsErr +} + func (m *mockLoginStore) GetActiveOrganizationOrDefault() (*entity.Organization, error) { m.getOrCreateOrgCalls++ return &entity.Organization{ID: "org-1", Name: "org"}, nil @@ -112,12 +123,12 @@ func (m *mockLoginStore) GetAllWorkspaces(_ *store.GetWorkspacesOptions) ([]enti func (m *mockLoginStore) GetCurrentWorkspaceID() (string, error) { return "", nil } func (m *mockLoginStore) GetWindowsDir() (string, error) { return "", nil } -func TestRunLoginWithAPIKey_SavesKeyAndOrgWithoutUserOrBackendOrgCalls(t *testing.T) { +func TestRunLoginWithAPIKey_SavesKeyAndResolvedOrg(t *testing.T) { auth := &mockLoginAuth{} - loginStore := &mockLoginStore{} + loginStore := &mockLoginStore{listOrgs: []entity.Organization{{ID: "org-test", Name: "TestOrg"}}} opts := LoginOptions{Auth: auth, LoginStore: loginStore} - err := opts.RunLogin(terminal.New(), "", " "+testAPIKey+" ", " org-test ", false, "", "") + err := opts.RunLogin(terminal.New(), "", " "+testAPIKey+" ", "", false, "", "") require.NoError(t, err) assert.Equal(t, 1, auth.apiKeyCalls) @@ -126,7 +137,7 @@ func TestRunLoginWithAPIKey_SavesKeyAndOrgWithoutUserOrBackendOrgCalls(t *testin assert.Equal(t, 1, loginStore.setDefaultOrgCalls) require.NotNil(t, loginStore.defaultOrg) assert.Equal(t, "org-test", loginStore.defaultOrg.ID) - assert.Equal(t, "org-test", loginStore.defaultOrg.Name) + assert.Equal(t, "TestOrg", loginStore.defaultOrg.Name) assert.Equal(t, 0, auth.tokenCalls) assert.Equal(t, 0, auth.loginCalls) assert.Equal(t, 0, loginStore.getCurrentUserCalls) @@ -165,7 +176,7 @@ func TestRunLoginWithAPIKey_RejectsConflictingFlags(t *testing.T) { func TestNewCmdLoginWithAPIKey_SkipsPostLoginHooks(t *testing.T) { auth := &mockLoginAuth{} - loginStore := &mockLoginStore{} + loginStore := &mockLoginStore{listOrgs: []entity.Organization{{ID: "org-test", Name: "TestOrg"}}} cmd := NewCmdLogin(terminal.New(), loginStore, auth) cmd.SetOut(&bytes.Buffer{}) cmd.SetErr(&bytes.Buffer{}) @@ -181,8 +192,24 @@ func TestNewCmdLoginWithAPIKey_SkipsPostLoginHooks(t *testing.T) { assert.Equal(t, 0, loginStore.userHomeDirCalls) } +func TestNewCmdLogin_OrgIDFlagDeprecationWarning(t *testing.T) { + auth := &mockLoginAuth{} + loginStore := &mockLoginStore{listOrgs: []entity.Organization{{ID: "org-test", Name: "TestOrg"}}} + cmd := NewCmdLogin(terminal.New(), loginStore, auth) + var out bytes.Buffer // cobra prints deprecated-flag warnings via c.Print -> OutOrStderr (stdout) + cmd.SetOut(&out) + cmd.SetErr(&bytes.Buffer{}) + cmd.SetArgs([]string{"--api-key", testAPIKey, "--org-id", "org-test"}) + + err := cmd.Execute() + + require.NoError(t, err) + assert.Contains(t, out.String(), "--org-id has been deprecated", "passing --org-id should warn") + assert.Contains(t, out.String(), "resolved automatically from the API key") +} + func TestNewCmdLogin_HidesAPIKeyFlagsFromHelp(t *testing.T) { - cmd := NewCmdLogin(terminal.New(), &mockLoginStore{}, &mockLoginAuth{}) + cmd := NewCmdLogin(terminal.New(), &mockLoginStore{listOrgs: []entity.Organization{{ID: "org-test", Name: "TestOrg"}}}, &mockLoginAuth{}) var out bytes.Buffer cmd.SetOut(&out) cmd.SetErr(&bytes.Buffer{}) @@ -195,28 +222,64 @@ func TestNewCmdLogin_HidesAPIKeyFlagsFromHelp(t *testing.T) { assert.NotContains(t, out.String(), "--org-id") } -func TestRunLoginWithAPIKey_RejectsMissingOrgID(t *testing.T) { - tests := []struct { - name string - apiKey string - orgID string - }{ - {name: "missing org id", apiKey: testAPIKey, orgID: " "}, - } +func TestRunLoginWithAPIKey_AutoResolvesOrgWhenOrgIDOmitted(t *testing.T) { + auth := &mockLoginAuth{} + loginStore := &mockLoginStore{listOrgs: []entity.Organization{{ID: "org-123", Name: "TestOrg"}}} + opts := LoginOptions{Auth: auth, LoginStore: loginStore} - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - auth := &mockLoginAuth{} - loginStore := &mockLoginStore{} - opts := LoginOptions{Auth: auth, LoginStore: loginStore} + err := opts.RunLogin(terminal.New(), "", testAPIKey, "", false, "", "") - err := opts.RunLogin(terminal.New(), "", tt.apiKey, tt.orgID, false, "", "") + require.NoError(t, err) + assert.Equal(t, 1, auth.apiKeyCalls) + assert.Equal(t, "org-123", auth.apiKeyOrgID, "resolved org ID should be saved") + require.NotNil(t, loginStore.defaultOrg) + assert.Equal(t, "org-123", loginStore.defaultOrg.ID) +} - require.Error(t, err) - assert.Equal(t, 0, auth.apiKeyCalls) - assert.Equal(t, 0, loginStore.setDefaultOrgCalls) - }) +func TestRunLoginWithAPIKey_ResolveOrgFailureRejects(t *testing.T) { + auth := &mockLoginAuth{} + loginStore := &mockLoginStore{listOrgsErr: assert.AnError} + opts := LoginOptions{Auth: auth, LoginStore: loginStore} + + err := opts.RunLogin(terminal.New(), "", testAPIKey, "", false, "", "") + + require.Error(t, err) + assert.Equal(t, 0, auth.apiKeyCalls, "must not save when the key can't be resolved/validated") + assert.Equal(t, 0, loginStore.setDefaultOrgCalls) +} + +func TestRunLoginWithAPIKey_FlagKeyActivatesOverEnvKey(t *testing.T) { + t.Setenv(authpkg.APIKeyEnvVar, authpkg.BrevAPIKeyPrefix+"env-key") + auth := &mockLoginAuth{} + var seenEnv string + loginStore := &mockLoginStore{} + loginStore.listOrgsErr = nil + loginStore.listOrgs = []entity.Organization{{ID: "org-flag", Name: "FlagOrg"}} + // Capture the env at ListOrganizations time to prove the flag key is active. + loginStore.listOrgsFn = func() ([]entity.Organization, error) { + seenEnv = os.Getenv(authpkg.APIKeyEnvVar) + return loginStore.listOrgs, nil } + opts := LoginOptions{Auth: auth, LoginStore: loginStore} + + err := opts.RunLogin(terminal.New(), "", testAPIKey, "", false, "", "") + + require.NoError(t, err) + assert.Equal(t, testAPIKey, seenEnv, "flag key must be active during org resolution") + assert.Equal(t, testAPIKey, auth.apiKey, "flag key must be persisted") + assert.Equal(t, "org-flag", auth.apiKeyOrgID, "flag key's org must be saved") +} + +func TestRunLogin_TokenLoginSuppressesEnvAPIKey(t *testing.T) { + t.Setenv(authpkg.APIKeyEnvVar, authpkg.BrevAPIKeyPrefix+"env-key") + auth := &mockLoginAuth{} + opts := LoginOptions{Auth: auth, LoginStore: &mockLoginStore{}} + + err := opts.RunLogin(terminal.New(), "some-login-token", "", "", false, "", "") + + require.NoError(t, err) + assert.Equal(t, "", os.Getenv(authpkg.APIKeyEnvVar), "browser/token login must clear BREV_API_KEY so the saved JWT is used") + assert.Equal(t, 0, auth.apiKeyCalls, "token login must not take the --api-key path") } func TestRunLoginWithOrgIDWithoutAPIKeyRejects(t *testing.T) { diff --git a/pkg/cmd/ls/ls.go b/pkg/cmd/ls/ls.go index 996897cff..cd4954bcd 100644 --- a/pkg/cmd/ls/ls.go +++ b/pkg/cmd/ls/ls.go @@ -132,7 +132,7 @@ func getOrgForRunLs(cliAuth auth.CLIAuth, lsStore LsStore, orgflag string) (*ent var org *entity.Organization if cliAuth.IsAPIKey() { if orgflag != "" { - return nil, breverrors.NewValidationError("api key auth is scoped to the org saved during login; --org is not supported") + return nil, breverrors.NewValidationError("api key auth is scoped to the org the key belongs to; --org is not supported") } org, err := lsStore.GetActiveOrganizationOrDefault() if err != nil { diff --git a/pkg/cmd/register/register.go b/pkg/cmd/register/register.go index 5e662a30f..c83de152b 100644 --- a/pkg/cmd/register/register.go +++ b/pkg/cmd/register/register.go @@ -5,6 +5,7 @@ import ( "context" "errors" "fmt" + "os" "strings" "time" @@ -12,6 +13,7 @@ import ( "connectrpc.com/connect" "github.com/google/uuid" + "github.com/brevdev/brev-cli/pkg/auth" "github.com/brevdev/brev-cli/pkg/config" "github.com/brevdev/brev-cli/pkg/entity" breverrors "github.com/brevdev/brev-cli/pkg/errors" @@ -77,8 +79,8 @@ func defaultRegisterDeps() registerDeps { } } -type OrgLister interface { - ListOrganizations() ([]entity.Organization, error) +func resolveAPIKey() string { + return strings.TrimSpace(os.Getenv(auth.APIKeyEnvVar)) } var ( @@ -88,9 +90,14 @@ This command registers this machine with Brev and brings up the Brev tunnel. Two modes are supported: • Interactive (default): run 'brev register' with no flags and follow prompts for device name and org. - • Non-interactive: use --name and --org. No prompts; both are required. - Use for scripts/CI. -` + • Non-interactive: use --name and --org. No prompts; --name is required, and + --org is required unless --api-key is supplied. Use for scripts/CI. + +Headless auth (credential chain): pass --api-key (a Brev API key) or set +the BREV_API_KEY environment variable to authenticate without the login +link; the key authenticates this register command only — run 'brev login +--api-key' afterward to stay logged in. If neither is set, the login-link +flow is used.` registerExample = ` # Interactive (prompts for device name, org, confirmations) brev register @@ -151,18 +158,34 @@ func runRegister(ctx context.Context, t *terminal.Terminal, s RegisterStore, opt return fmt.Errorf("sudo issue: %w", err) } + apiKey := resolveAPIKey() + if apiKey != "" { + if !auth.IsBrevAPIKey(apiKey) { + return breverrors.NewValidationError(fmt.Sprintf("api key must be a Brev API key (expected %s prefix); see 'brev login --api-key'", auth.BrevAPIKeyPrefix)) + } + } if !opts.interactive { - if opts.name == "" || opts.orgName == "" { - return fmt.Errorf("in non-interactive mode --name and --org are required") + if opts.name == "" { + return fmt.Errorf("in non-interactive mode --name is required") + } + if opts.orgName == "" && apiKey == "" { + return fmt.Errorf("in non-interactive mode --org is required unless --api-key is supplied") } } - // Verify the user is authenticated before performing any local side effects. - if _, err := s.GetCurrentUser(); err != nil { + + if err := isAuthenticated(s, apiKey); err != nil { return breverrors.WrapAndTrace(err) } var intendedOrg *entity.Organization - if !opts.interactive { + switch { + case apiKey != "": + o, err := ResolveOrgForAPIKey(s, opts.orgName) + if err != nil { + return err + } + intendedOrg = o + case !opts.interactive: o, err := resolveOrg(s, opts.orgName) if err != nil { return err @@ -349,6 +372,16 @@ func resolveOrgInteractive(t *terminal.Terminal, s RegisterStore, deps registerD return org, nil } +func isAuthenticated(s RegisterStore, apiKey string) error { + if apiKey != "" { + return nil + } + if _, err := s.GetCurrentUser(); err != nil { + return breverrors.WrapAndTrace(err) + } + return nil +} + func resolveOrg(s RegisterStore, orgName string) (*entity.Organization, error) { org, err := helpers.ResolveOrgByName(s, orgName) if err != nil { @@ -357,6 +390,21 @@ func resolveOrg(s RegisterStore, orgName string) (*entity.Organization, error) { return org, nil } +func ResolveOrgForAPIKey(s auth.OrgLister, orgName string) (*entity.Organization, error) { + orgs, err := s.ListOrganizations() + if err != nil { + return nil, breverrors.WrapAndTrace(err) + } + org, err := auth.SingleOrgForAPIKey(orgs) + if err != nil { + return nil, breverrors.WrapAndTrace(err) + } + if orgName != "" && org.Name != orgName { + return nil, breverrors.NewValidationError(fmt.Sprintf("api key does not belong to organization %q", orgName)) + } + return org, nil +} + func orgMismatchError(reg *DeviceRegistration, intended *entity.Organization) error { existing := "this device is already registered in org" if reg.Status == RegistrationStatusPending { diff --git a/pkg/cmd/register/register_test.go b/pkg/cmd/register/register_test.go index 9e1f60d0f..0124bfb2b 100644 --- a/pkg/cmd/register/register_test.go +++ b/pkg/cmd/register/register_test.go @@ -12,6 +12,7 @@ import ( nodev1 "buf.build/gen/go/brevdev/devplane/protocolbuffers/go/devplaneapi/v1" "connectrpc.com/connect" + "github.com/brevdev/brev-cli/pkg/auth" "github.com/brevdev/brev-cli/pkg/entity" "github.com/brevdev/brev-cli/pkg/externalnode" "github.com/brevdev/brev-cli/pkg/sudo" @@ -214,71 +215,72 @@ func testPendingReg(orgID, orgName, deviceID string) *DeviceRegistration { } } -func Test_runRegister_HappyPath(t *testing.T) { - regStore := &mockRegistrationStore{} - - store := testRegisterStore() +func registeredReg() *DeviceRegistration { + return &DeviceRegistration{ + ExternalNodeID: "unode_existing", + DisplayName: "Existing", + OrgID: "org_123", + DeviceID: "dev-existing", + Status: RegistrationStatusRegistered, + } +} - svc := &fakeNodeService{ - addNodeFn: func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { - if req.GetOrganizationId() != "org_123" { - t.Errorf("unexpected org: %s", req.GetOrganizationId()) - } - if req.GetName() != "my-spark" { - t.Errorf("unexpected name: %s", req.GetName()) - } - return &nodev1.AddNodeResponse{ - ExternalNode: &nodev1.ExternalNode{ - ExternalNodeId: "unode_abc", - OrganizationId: "org_123", - Name: req.GetName(), - DeviceId: req.GetDeviceId(), - ConnectivityInfo: &nodev1.ConnectivityInfo{ - RegistrationCommand: "netbird up --key abc", - }, +// okAddNodeFn echoes org/name/device and returns the standard success node. +func okAddNodeFn(captured *[]string) func(*nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { + return func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { + if captured != nil { + *captured = append(*captured, req.GetOrganizationId()) + } + return &nodev1.AddNodeResponse{ + ExternalNode: &nodev1.ExternalNode{ + ExternalNodeId: "unode_abc", + OrganizationId: req.GetOrganizationId(), + Name: req.GetName(), + DeviceId: req.GetDeviceId(), + ConnectivityInfo: &nodev1.ConnectivityInfo{ + RegistrationCommand: "netbird up --key abc", }, - }, nil - }, + }, + }, nil } +} - setupRunner := &mockSetupRunner{} - +// runRegisterCase absorbs scaffolding tests repeat: deps, server +// lifecycle, terminal, and the invocation on the standard test store. +func runRegisterCase(t *testing.T, regStore RegistrationStore, svc *fakeNodeService, opts registerOpts) error { + t.Helper() deps, server := testRegisterDeps(t, svc, regStore) defer server.Close() + return runRegister(context.Background(), terminal.New(), testRegisterStore(), opts, deps) +} + +func Test_runRegister_HappyPath(t *testing.T) { + regStore := &mockRegistrationStore{} + var gotOrgs []string + svc := &fakeNodeService{addNodeFn: okAddNodeFn(&gotOrgs)} + setupRunner := &mockSetupRunner{} + deps, server := testRegisterDeps(t, svc, regStore) + defer server.Close() deps.setupRunner = setupRunner - term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg"} - err := runRegister(context.Background(), term, store, opts, deps) + err := runRegister(context.Background(), terminal.New(), testRegisterStore(), + registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg"}, deps) if err != nil { t.Fatalf("runRegister failed: %v", err) } - // Verify registration was persisted - exists, err := regStore.Exists() - if err != nil { - t.Fatalf("Exists error: %v", err) + if len(gotOrgs) != 1 || gotOrgs[0] != "org_123" { + t.Errorf("expected AddNode for org_123 once, got %v", gotOrgs) } - if !exists { - t.Fatal("expected registration to exist after successful register") - } - - reg, err := regStore.Load() + reg, err := regStore.Load() // implies Exists if err != nil { t.Fatalf("Load failed: %v", err) } - if reg.ExternalNodeID != "unode_abc" { - t.Errorf("expected ExternalNodeID unode_abc, got %s", reg.ExternalNodeID) - } - if reg.DisplayName != "my-spark" { - t.Errorf("expected display name 'my-spark', got %s", reg.DisplayName) + if reg.ExternalNodeID != "unode_abc" || reg.DisplayName != "my-spark" || + reg.OrgID != "org_123" || reg.Status != RegistrationStatusRegistered { + t.Errorf("persisted registration mismatch: %+v", reg) } - if reg.OrgID != "org_123" { - t.Errorf("expected org org_123, got %s", reg.OrgID) - } - - // Verify setup command was executed if setupRunner.cmd != "netbird up --key abc" { t.Errorf("expected setup command 'netbird up --key abc', got %q", setupRunner.cmd) } @@ -379,13 +381,7 @@ func Test_runRegister_AlreadyRegistered(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - regStore := &mockRegistrationStore{ - reg: &DeviceRegistration{ - ExternalNodeID: "unode_existing", - DisplayName: "Existing", - OrgID: "org_123", - }, - } + regStore := &mockRegistrationStore{reg: registeredReg()} store := testRegisterStore() @@ -442,20 +438,8 @@ func Test_runRegister_WithOrgFlag(t *testing.T) { token: "tok", } - var capturedOrgID string - svc := &fakeNodeService{ - addNodeFn: func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { - capturedOrgID = req.GetOrganizationId() - return &nodev1.AddNodeResponse{ - ExternalNode: &nodev1.ExternalNode{ - ExternalNodeId: "unode_abc", - OrganizationId: req.GetOrganizationId(), - Name: req.GetName(), - DeviceId: req.GetDeviceId(), - }, - }, nil - }, - } + var gotOrgs []string + svc := &fakeNodeService{addNodeFn: okAddNodeFn(&gotOrgs)} setupRunner := &mockSetupRunner{} deps, server := testRegisterDeps(t, svc, regStore) @@ -466,20 +450,16 @@ func Test_runRegister_WithOrgFlag(t *testing.T) { opts := registerOpts{interactive: false, name: "my-spark", orgName: tt.orgName} err := runRegister(context.Background(), term, store, opts, deps) if tt.wantErr != "" { - if err == nil { - t.Fatal("expected error when org not found") - } - if !strings.Contains(err.Error(), tt.wantErr) { - t.Errorf("expected %q error, got: %v", tt.wantErr, err) + if err == nil || !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("expected error containing %q, got: %v", tt.wantErr, err) } return } if err != nil { t.Fatalf("runRegister with --org failed: %v", err) } - - if capturedOrgID != tt.wantOrgID { - t.Errorf("expected org %s, got %s", tt.wantOrgID, capturedOrgID) + if len(gotOrgs) != 1 || gotOrgs[0] != tt.wantOrgID { + t.Errorf("expected AddNode org %s, got %v", tt.wantOrgID, gotOrgs) } reg, err := regStore.Load() @@ -507,19 +487,13 @@ func Test_runRegister_AddNodeFailure(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { regStore := &mockRegistrationStore{} - store := testRegisterStore() svc := &fakeNodeService{ addNodeFn: func(_ *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { return nil, connect.NewError(tt.code, errors.New(tt.errMsg)) }, } - deps, server := testRegisterDeps(t, svc, regStore) - defer server.Close() - - term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg"} - err := runRegister(context.Background(), term, store, opts, deps) + err := runRegisterCase(t, regStore, svc, registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg"}) if err == nil { t.Fatal("expected error on AddNode failure") } @@ -534,15 +508,12 @@ func Test_runRegister_AddNodeFailure(t *testing.T) { if !tt.wantPending { return } - reg, loadErr := regStore.LoadAll() - if loadErr != nil { - t.Fatalf("Load failed: %v", loadErr) - } - if reg.Status != RegistrationStatusPending { - t.Errorf("expected pending status, got %q", reg.Status) + reg, err := regStore.LoadAll() + if err != nil { + t.Fatalf("Load failed: %v", err) } - if reg.DeviceID == "" { - t.Error("expected pending record to carry a device ID for retry") + if reg.Status != RegistrationStatusPending || reg.DeviceID == "" { + t.Errorf("pending record must carry a device ID for retry: status=%q device=%q", reg.Status, reg.DeviceID) } }) } @@ -792,28 +763,111 @@ func Test_runRegister_ResumesPendingRegistration(t *testing.T) { } } -// --- Org mismatch --- +// --- API key org scoping --- + +const testAPIKey = auth.BrevAPIKeyPrefix + "test-key" + +func ensureNoAPIKeyEnv(t *testing.T) { + t.Helper() + t.Setenv(auth.APIKeyEnvVar, "") +} + +func Test_resolveOrgForAPIKey(t *testing.T) { + tests := []struct { + name string + orgs []entity.Organization + orgName string + wantID string + wantErr string + }{ + {"single, name matches", []entity.Organization{{ID: "org_1", Name: "Alpha"}}, "Alpha", "org_1", ""}, + {"single, no name", []entity.Organization{{ID: "org_1", Name: "Alpha"}}, "", "org_1", ""}, + {"single, name mismatch", []entity.Organization{{ID: "org_1", Name: "Alpha"}}, "Beta", "", "does not belong to organization"}, + {"empty", nil, "", "", "api key invalid"}, + {"multiple, name matches one", []entity.Organization{{ID: "org_1", Name: "Alpha"}, {ID: "org_2", Name: "Beta"}}, "Beta", "", "api key invalid"}, + {"multiple, no name", []entity.Organization{{ID: "org_1", Name: "Alpha"}, {ID: "org_2", Name: "Beta"}}, "", "", "api key invalid"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + s := &mockRegisterStore{orgs: tt.orgs} + got, err := ResolveOrgForAPIKey(s, tt.orgName) + if tt.wantErr != "" { + if err == nil { + t.Fatalf("expected error containing %q, got nil", tt.wantErr) + } + if !strings.Contains(err.Error(), tt.wantErr) { + t.Errorf("expected error containing %q, got: %v", tt.wantErr, err) + } + return + } + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got.ID != tt.wantID { + t.Errorf("expected org ID %s, got %s", tt.wantID, got.ID) + } + }) + } +} + +// API-key flows: mismatch rejects with guidance; no --org uses the key's org. +func Test_runRegister_APIKey(t *testing.T) { + t.Setenv(auth.APIKeyEnvVar, testAPIKey) + + tests := []struct { + name string + orgName string + wantErr string + wantOrgID string // AddNode must receive this org + }{ + {"org flag mismatch rejects", "OtherOrg", "does not belong to organization", ""}, + {"no org flag uses key org", "", "", "org_123"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + regStore := &mockRegistrationStore{} + var gotOrgs []string + svc := &fakeNodeService{addNodeFn: okAddNodeFn(&gotOrgs)} + + err := runRegisterCase(t, regStore, svc, registerOpts{interactive: false, name: "my-spark", orgName: tt.orgName}) + + if tt.wantErr != "" { + if err == nil || !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("expected error containing %q, got: %v", tt.wantErr, err) + } + return + } + if err != nil { + t.Fatalf("runRegister failed: %v", err) + } + if len(gotOrgs) != 1 || gotOrgs[0] != tt.wantOrgID { + t.Errorf("expected AddNode org %s, got %v", tt.wantOrgID, gotOrgs) + } + }) + } +} func Test_runRegister_OrgMismatch(t *testing.T) { tests := []struct { name string status string + useAPIKey bool // set BREV_API_KEY for the new org useOrgFlag bool // pass --org for the new org wantWording string // pending -> "incomplete registration"; registered -> "already registered" }{ - {"OrgFlag_Pending", RegistrationStatusPending, true, "incomplete registration"}, - {"OrgFlag_AlreadyRegistered", RegistrationStatusRegistered, true, "already registered"}, + {"APIKey_Pending", RegistrationStatusPending, true, false, "incomplete registration"}, + {"OrgFlag_Pending", RegistrationStatusPending, false, true, "incomplete registration"}, + {"APIKey_AlreadyRegistered", RegistrationStatusRegistered, true, false, "already registered"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - regStore := &mockRegistrationStore{reg: &DeviceRegistration{ - ExternalNodeID: "unode_existing", - DisplayName: "My Spark", - OrgID: "org_other", - OrgName: "OtherOrg", - DeviceID: "dev-pending", - Status: tt.status, - }} + ensureNoAPIKeyEnv(t) + if tt.useAPIKey { + t.Setenv(auth.APIKeyEnvVar, testAPIKey) + } + reg := registeredReg() + reg.OrgID, reg.OrgName, reg.Status = "org_other", "OtherOrg", tt.status + regStore := &mockRegistrationStore{reg: reg} store := &mockRegisterStore{ user: &entity.User{ID: "user_1"}, orgs: []entity.Organization{{ID: "org_123", Name: "TestOrg"}}, // new org @@ -837,14 +891,10 @@ func Test_runRegister_OrgMismatch(t *testing.T) { if err == nil { t.Fatal("expected error on org mismatch") } - if !strings.Contains(err.Error(), "deregister") { - t.Errorf("expected deregister guidance, got: %v", err) - } - if !strings.Contains(err.Error(), tt.wantWording) { - t.Errorf("expected %q wording, got: %v", tt.wantWording, err) - } - if !strings.Contains(err.Error(), "org_other") || !strings.Contains(err.Error(), "org_123") { - t.Errorf("expected both org IDs in message, got: %v", err) + for _, want := range []string{"deregister", tt.wantWording, "org_other", "org_123"} { + if !strings.Contains(err.Error(), want) { + t.Errorf("expected error to contain %q, got: %v", want, err) + } } if addNodeCalls != 0 { t.Errorf("AddNode must not be called on mismatch, got %d", addNodeCalls) @@ -857,8 +907,6 @@ func Test_runRegister_ResumeAddNodeFails_StaysPending(t *testing.T) { const pendingDeviceID = "device-uuid-pending" regStore := &mockRegistrationStore{reg: testPendingReg("org_123", "TestOrg", pendingDeviceID)} - store := testRegisterStore() - var addNodeCalls int svc := &fakeNodeService{ addNodeFn: func(_ *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { @@ -867,11 +915,7 @@ func Test_runRegister_ResumeAddNodeFails_StaysPending(t *testing.T) { }, } - deps, server := testRegisterDeps(t, svc, regStore) - defer server.Close() - - term := terminal.New() - err := runRegister(context.Background(), term, store, registerOpts{interactive: true}, deps) + err := runRegisterCase(t, regStore, svc, registerOpts{interactive: true}) if err == nil { t.Fatal("expected error when AddNode fails during resume") } @@ -879,17 +923,41 @@ func Test_runRegister_ResumeAddNodeFails_StaysPending(t *testing.T) { t.Fatalf("expected AddNode called once, got %d", addNodeCalls) } - reg, loadErr := regStore.LoadAll() - if loadErr != nil { - t.Fatalf("Load failed: %v", loadErr) - } - if reg.Status != RegistrationStatusPending { - t.Errorf("expected record to stay pending after failed resume, got %q", reg.Status) - } - if reg.DeviceID != pendingDeviceID { - t.Errorf("expected device ID to remain %q, got %q", pendingDeviceID, reg.DeviceID) + reg, err := regStore.LoadAll() + if err != nil { + t.Fatalf("Load failed: %v", err) } - if reg.ExternalNodeID != "" { - t.Errorf("expected no ExternalNodeID after failed resume, got %q", reg.ExternalNodeID) + if reg.Status != RegistrationStatusPending || reg.DeviceID != pendingDeviceID || reg.ExternalNodeID != "" { + t.Errorf("pending record must survive with device ID intact: status=%q device=%q node=%q", + reg.Status, reg.DeviceID, reg.ExternalNodeID) } } + +func Test_isAuthenticated(t *testing.T) { + t.Run("api key short-circuits without GetCurrentUser", func(t *testing.T) { + store := &mockRegisterStore{ + err: fmt.Errorf("GetCurrentUser must not be called when a key is present"), + } + if err := isAuthenticated(store, testAPIKey); err != nil { + t.Fatalf("expected nil error with api key, got %v", err) + } + }) + + t.Run("no api key verifies via GetCurrentUser", func(t *testing.T) { + store := &mockRegisterStore{user: &entity.User{ID: "user_1"}} + if err := isAuthenticated(store, ""); err != nil { + t.Fatalf("expected nil error with valid user, got %v", err) + } + }) + + t.Run("no api key and GetCurrentUser fails", func(t *testing.T) { + store := &mockRegisterStore{err: fmt.Errorf("not logged in")} + err := isAuthenticated(store, "") + if err == nil { + t.Fatal("expected error when GetCurrentUser fails") + } + if !strings.Contains(err.Error(), "not logged in") { + t.Errorf("expected GetCurrentUser error propagated, got %v", err) + } + }) +} diff --git a/pkg/errors/errors.go b/pkg/errors/errors.go index 236a0f7aa..6579aacea 100644 --- a/pkg/errors/errors.go +++ b/pkg/errors/errors.go @@ -125,9 +125,11 @@ func (v ValidationError) Error() string { type DeclineToLoginError struct{} -func (d *DeclineToLoginError) Error() string { return "declined to login" } +func (d *DeclineToLoginError) Error() string { return DeclineToLoginMessage } func (d *DeclineToLoginError) Directive() string { return "log in to run this command" } +const DeclineToLoginMessage = "declined to login" + var NetworkErrorMessage = "possible internet connection problem" type CredentialsFileNotFound struct{} diff --git a/pkg/store/http.go b/pkg/store/http.go index 60884f810..7fc64f61b 100644 --- a/pkg/store/http.go +++ b/pkg/store/http.go @@ -3,13 +3,16 @@ package store import ( "encoding/json" "fmt" + "log" "net/http" + "os" "runtime" "strings" "github.com/brevdev/brev-cli/pkg/cmd/version" breverrors "github.com/brevdev/brev-cli/pkg/errors" "github.com/brevdev/brev-cli/pkg/featureflag" + resty "github.com/go-resty/resty/v2" ) @@ -151,6 +154,64 @@ func WithDebug(debug bool) Option { } } +// quietRestyLogger swallows resty's retry WARN/ERROR chatter for expected, +// user-driven errors (declined login). Everything else logs as before. +type quietRestyLogger struct { + next resty.Logger +} + +func (q quietRestyLogger) Errorf(format string, v ...interface{}) { + if isDeclinedLoginMsg(format, v...) { + return + } + q.next.Errorf(format, v...) +} + +func (q quietRestyLogger) Warnf(format string, v ...interface{}) { + if isDeclinedLoginMsg(format, v...) { + return + } + q.next.Warnf(format, v...) +} + +func (q quietRestyLogger) Debugf(format string, v ...interface{}) { + q.next.Debugf(format, v...) +} + +func isDeclinedLoginMsg(format string, v ...interface{}) bool { + if !strings.Contains(format, "%v") { // formatted messages embed the error + return false + } + msg := fmt.Sprintf(format, v...) + return strings.Contains(msg, breverrors.DeclineToLoginMessage) +} + +// stderrLogger mirrors resty's default logger (stderr, date+microseconds, +// "WARN RESTY"/"ERROR RESTY" prefixes) so quietRestyLogger has a real sink. +type stderrLogger struct { + l *log.Logger +} + +func newStderrLogger() stderrLogger { + return stderrLogger{l: log.New(os.Stderr, "", log.Ldate|log.Lmicroseconds)} +} + +func (s stderrLogger) Errorf(format string, v ...interface{}) { + s.outputf("ERROR RESTY "+format, v...) +} + +func (s stderrLogger) Warnf(format string, v ...interface{}) { + s.outputf("WARN RESTY "+format, v...) +} + +func (s stderrLogger) Debugf(format string, v ...interface{}) { + s.outputf("DEBUG RESTY "+format, v...) +} + +func (s stderrLogger) outputf(format string, v ...interface{}) { + _ = s.l.Output(2, fmt.Sprintf(format, v...)) +} + func NewAuthHTTPClient(auth Auth, brevAPIURL string, options ...Option) *AuthHTTPClient { opts := &Options{} for _, o := range options { @@ -158,6 +219,10 @@ func NewAuthHTTPClient(auth Auth, brevAPIURL string, options ...Option) *AuthHTT } restyClient := NewRestyClient(brevAPIURL) restyClient.Debug = opts.Debug + // quietRestyLogger wraps a real stderr logger (matching resty's default + // format) and swallows only declined-login retry chatter. Everything else + // — genuine HTTP errors, debug output when Debug is on — still logs. + restyClient.SetLogger(quietRestyLogger{next: newStderrLogger()}) restyClient.OnBeforeRequest(func(c *resty.Client, r *resty.Request) error { token, err := auth.GetAccessToken() if err != nil { diff --git a/pkg/store/http_test.go b/pkg/store/http_test.go index 961ee2c4c..ca0fdf3f5 100644 --- a/pkg/store/http_test.go +++ b/pkg/store/http_test.go @@ -1,35 +1,33 @@ package store import ( - "net/http" - "strings" + "bytes" + "errors" + stderrors "errors" + "fmt" + "io" + "os" "testing" + "time" - "github.com/jarcoal/httpmock" + breverrors "github.com/brevdev/brev-cli/pkg/errors" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -func MakeMockNoHTTPStore() *NoAuthHTTPStore { - fs := MakeMockFileStore() - nh := fs.WithNoAuthHTTPClient(NewNoAuthHTTPClient("")) - return nh -} - -func TestWithHTTPClient(t *testing.T) { - nh := MakeMockNoHTTPStore() - if !assert.NotNil(t, nh) { - return - } -} - type MockAuth struct{ token *string } func (a MockAuth) GetAccessToken() (string, error) { if a.token == nil { return "mock-token", nil - } else { - return *a.token, nil } + return *a.token, nil +} + +func MakeMockNoHTTPStore() *NoAuthHTTPStore { + fs := MakeMockFileStore() + nh := fs.WithNoAuthHTTPClient(NewNoAuthHTTPClient("")) + return nh } func MakeMockAuthHTTPStore() *AuthHTTPStore { @@ -38,96 +36,128 @@ func MakeMockAuthHTTPStore() *AuthHTTPStore { return ah } -func TestWithAuthHTTPClient(t *testing.T) { - ah := MakeMockAuthHTTPStore() - if !assert.NotNil(t, ah) { - return - } +func TestQuietRestyLogger_SuppressesDeclineLogin(t *testing.T) { + var buf bytes.Buffer + base := &testLogger{out: &buf} + q := quietRestyLogger{next: base} + + q.Errorf("%v", errors.New(breverrors.DeclineToLoginMessage)) + q.Warnf("%v, Attempt %v", errors.New(breverrors.DeclineToLoginMessage), 1) + q.Errorf("%v", errors.New("connection refused")) + q.Warnf("some other warning") + + out := buf.String() + assert.NotContains(t, out, "declined to login") + assert.Contains(t, out, "connection refused") + assert.Contains(t, out, "some other warning") } -func TestNewNoAuthHTTPClient(t *testing.T) { - n := NewNoAuthHTTPClient("") - if !assert.NotNil(t, n) { - return - } +func TestQuietRestyLogger_DebugPassesThrough(t *testing.T) { + var buf bytes.Buffer + base := &testLogger{out: &buf} + q := quietRestyLogger{next: base} + + q.Debugf("debug %s", "detail") + assert.Contains(t, buf.String(), "debug detail") } -func TestNewAuthHTTPClient(t *testing.T) { - n := NewAuthHTTPClient(MockAuth{}, "") - if !assert.NotNil(t, n) { - return - } +func TestIsDeclinedLoginMsg(t *testing.T) { + assert.True(t, isDeclinedLoginMsg("%v", errors.New("declined to login"))) + assert.False(t, isDeclinedLoginMsg("%v", errors.New("boom"))) + // Non-%v formats carry no embedded error; never filtered. + assert.False(t, isDeclinedLoginMsg("plain format")) } -func makeCheckTokenResponder(validToken string) httpmock.Responder { - return func(r *http.Request) (*http.Response, error) { - h := r.Header.Get("Authorization") - splitStr := strings.Split(h, "Bearer") - if len(splitStr) != 2 { - return &http.Response{StatusCode: 403}, nil - } - if strings.TrimSpace(splitStr[1]) == validToken { - return &http.Response{StatusCode: 200}, nil - } else { - return &http.Response{StatusCode: 403}, nil - } - } +type testLogger struct { + out *bytes.Buffer } -func TestRetryAuthSuccess(t *testing.T) { - nh := MakeMockNoHTTPStore() - invalidToken := "invalid-token" - validToken := "valid-token" - s := nh.WithAuthHTTPClient(NewAuthHTTPClient(MockAuth{&invalidToken}, "")) +func (t *testLogger) Errorf(format string, v ...interface{}) { t.writef(format, v...) } +func (t *testLogger) Warnf(format string, v ...interface{}) { t.writef(format, v...) } +func (t *testLogger) Debugf(format string, v ...interface{}) { t.writef(format, v...) } - httpmock.ActivateNonDefault(s.authHTTPClient.restyClient.GetClient()) +func (t *testLogger) writef(format string, v ...interface{}) { + fmt.Fprintf(t.out, format, v...) +} - url := "/test" - responder := makeCheckTokenResponder(validToken) - httpmock.RegisterResponder("GET", url, responder) +// declineAuth simulates a user answering "n" at the login prompt. +type declineAuth struct{} - calledTimes := 0 - err := s.SetForbiddenStatusRetryHandler(func() error { - invalidToken = validToken - calledTimes++ - return nil - }) - if !assert.Nil(t, err) { - return - } - r, err := s.authHTTPClient.restyClient.R().Get(url) - if !assert.Nil(t, err) { - return - } - if !assert.Equal(t, 200, r.StatusCode()) { - return - } - if !assert.Equal(t, 1, calledTimes) { - return - } +func (declineAuth) GetAccessToken() (string, error) { + return "", &breverrors.DeclineToLoginError{} } -func TestRetryAuthFailure(t *testing.T) { - s := MakeMockAuthHTTPStore() - httpmock.ActivateNonDefault(s.authHTTPClient.restyClient.GetClient()) +// The exact scenario from the bug report: a command runs, the user declines +// login, and the request fails. Resty must NOT spray WARN/ERROR retry chatter +// with stack-traced wrappers to stderr; the error must surface as a clean +// sentinel for DisplayAndHandleError to render. +func TestNewAuthHTTPClient_DeclinedLoginIsQuietAndClean(t *testing.T) { + // NO sink replacement: the factory-installed logger chain must handle + // this itself. Capture stderr (where the factory's logger writes) and + // assert the decline chatter never reaches it. + r, w, err := os.Pipe() + require.NoError(t, err) + origStderr := os.Stderr + os.Stderr = w + t.Cleanup(func() { + os.Stderr = origStderr + _ = r.Close() + _ = w.Close() + }) + + client := NewAuthHTTPClient(declineAuth{}, "https://api.test") // installs quietRestyLogger{next: stderrLogger} + client.restyClient.SetRetryCount(1) + client.restyClient.SetTimeout(2 * time.Second) + + _, err = client.restyClient.R().Get("/user") + + require.NoError(t, w.Close()) + var buf bytes.Buffer + _, _ = io.Copy(&buf, r) + os.Stderr = origStderr - url := "/test" - res := httpmock.NewStringResponder(403, "") - httpmock.RegisterResponder("GET", url, res) + require.Error(t, err) + var decline *breverrors.DeclineToLoginError + require.True(t, breverrors.As(err, &decline), "decline sentinel must survive the resty round trip (wrapping allowed)") + require.True(t, stderrors.Is(err, decline), "DisplayAndHandleError matches the sentinel with errors.Is through wrapping") + assert.NotContains(t, buf.String(), "declined to login", "factory logger must suppress decline retry chatter") +} - calledTimes := 0 - err := s.SetForbiddenStatusRetryHandler(func() error { - calledTimes++ - return nil +// The factory must install quietRestyLogger over a REAL sink: unrelated +// errors still reach stderr (only declined-login chatter is filtered). +func TestNewAuthHTTPClient_LoggerForwardsUnrelatedErrors(t *testing.T) { + // Replace stderr before construction so the factory's stderrLogger + // captures our pipe. + r, w, err := os.Pipe() + require.NoError(t, err) + origStderr := os.Stderr + os.Stderr = w + t.Cleanup(func() { + os.Stderr = origStderr + _ = r.Close() + _ = w.Close() }) - if !assert.Nil(t, err) { - return - } - _, err = s.authHTTPClient.restyClient.R().Get(url) - if !assert.Nil(t, err) { - return - } - if !assert.Equal(t, 1, calledTimes) { - return - } + + client := NewAuthHTTPClient(errorAuth{}, "https://api.test") + client.restyClient.SetRetryCount(1) + client.restyClient.SetTimeout(2 * time.Second) + + _, err = client.restyClient.R().Get("/user") + + // Flush the pipe before restoring. + _ = w.Close() + var buf bytes.Buffer + _, _ = io.Copy(&buf, r) + os.Stderr = origStderr + + require.Error(t, err) + assert.Contains(t, buf.String(), "ERROR RESTY", "unrelated auth errors must still be logged by the factory logger") + assert.Contains(t, buf.String(), "boom-auth", "the actual error text must reach the sink") +} + +// errorAuth fails auth with a non-decline error: must be loud. +type errorAuth struct{} + +func (errorAuth) GetAccessToken() (string, error) { + return "", errors.New("boom-auth") } diff --git a/pkg/store/organization.go b/pkg/store/organization.go index bc02aedb4..c47453f5c 100644 --- a/pkg/store/organization.go +++ b/pkg/store/organization.go @@ -72,26 +72,7 @@ func (f FileStore) GetCachedActiveOrganizationOrNil() (*entity.Organization, err // returns the 'set'/active organization or nil if not set func (s AuthHTTPStore) GetActiveOrganizationOrNil() (*entity.Organization, error) { if auth.IsAPIKeyAuthStore(&s) { - orgID, err := auth.GetAPIKeyOrgID(&s) - if err != nil { - return nil, breverrors.WrapAndTrace(err) - } - org := &entity.Organization{ID: orgID, Name: orgID} - // Name hydration is best-effort; the command itself should surface backend auth errors. - freshOrg, err := s.GetOrganization(orgID) - if err != nil { - return org, nil - } - if freshOrg == nil { - return org, nil - } - if freshOrg.ID == "" { - freshOrg.ID = orgID - } - if freshOrg.Name == "" { - freshOrg.Name = freshOrg.ID - } - return freshOrg, nil + return s.hydrateOrgFromApiKey() } workspaceID, err := s.GetCurrentWorkspaceID() @@ -99,17 +80,7 @@ func (s AuthHTTPStore) GetActiveOrganizationOrNil() (*entity.Organization, error return nil, breverrors.WrapAndTrace(err) } if workspaceID != "" { - var workspace *entity.Workspace - workspace, err = s.GetWorkspace(workspaceID) - if err != nil { - return nil, breverrors.WrapAndTrace(err) - } - var org *entity.Organization - org, err = s.GetOrganization(workspace.OrganizationID) - if err != nil { - return nil, breverrors.WrapAndTrace(err) - } - return org, nil + return s.hydrateOrgFromWorkspace(workspaceID) } activeOrg, err := s.GetCachedActiveOrganizationOrNil() @@ -129,6 +100,50 @@ func (s AuthHTTPStore) GetActiveOrganizationOrNil() (*entity.Organization, error return freshOrg, nil } +func (s AuthHTTPStore) hydrateOrgFromWorkspace(workspaceID string) (*entity.Organization, error) { + var workspace *entity.Workspace + workspace, err := s.GetWorkspace(workspaceID) + if err != nil { + return nil, breverrors.WrapAndTrace(err) + } + var org *entity.Organization + org, err = s.GetOrganization(workspace.OrganizationID) + if err != nil { + return nil, breverrors.WrapAndTrace(err) + } + return org, nil +} + +func (s AuthHTTPStore) hydrateOrgFromApiKey() (*entity.Organization, error) { + org, err := auth.ResolveEnvAPIKeyOrg(&s) + if err != nil { + return nil, breverrors.WrapAndTrace(err) + } + if org != nil { + return org, nil + } + orgID, err := auth.GetAPIKeyOrgID(&s) + if err != nil { + return nil, breverrors.WrapAndTrace(err) + } + org = &entity.Organization{ID: orgID, Name: orgID} + // Name hydration is best-effort; the command itself should surface backend auth errors. + freshOrg, err := s.GetOrganization(orgID) + if err != nil { + return org, nil + } + if freshOrg == nil { + return org, nil + } + if freshOrg.ID == "" { + freshOrg.ID = orgID + } + if freshOrg.Name == "" { + freshOrg.Name = freshOrg.ID + } + return freshOrg, nil +} + // returns the 'set'/active organization or the default one or nil if no orgs exist func (s AuthHTTPStore) GetActiveOrganizationOrDefault() (*entity.Organization, error) { org, err := s.GetActiveOrganizationOrNil() diff --git a/pkg/store/organization_test.go b/pkg/store/organization_test.go index ec7f0cebb..db3b0ca2b 100644 --- a/pkg/store/organization_test.go +++ b/pkg/store/organization_test.go @@ -162,6 +162,33 @@ func TestGetActiveOrganization_APIKeyUsesCredentialOrgNameWhenAvailable(t *testi assert.Equal(t, expected.Name, org.Name) } +func TestGetActiveOrganization_APIKeyEnvResolvesOrgInRealTime(t *testing.T) { + apiKey := authpkg.BrevAPIKeyPrefix + "env-key" + fileStore, _, _ := newAuthTokenTestStore(t) + s := fileStore.WithAuthHTTPClient(NewAuthHTTPClient(MockAuth{token: &apiKey}, "https://api.test")) + httpmock.ActivateNonDefault(s.authHTTPClient.restyClient.GetClient()) + defer httpmock.DeactivateAndReset() + + require.NoError(t, s.SaveAuthTokens(entity.AuthTokens{ + APIKey: authpkg.BrevAPIKeyPrefix + "other-key", + APIKeyOrgID: "org-stale", + })) + t.Setenv(authpkg.APIKeyEnvVar, apiKey) + + expected := []entity.Organization{{ID: "org-real", Name: "Real Org"}} + res, err := httpmock.NewJsonResponder(200, expected) + require.NoError(t, err) + url := fmt.Sprintf("%s/%s", s.authHTTPClient.restyClient.BaseURL, orgPath) + httpmock.RegisterResponder("GET", url, res) + + org, err := s.GetActiveOrganizationOrDefault() + + require.NoError(t, err) + require.NotNil(t, org) + assert.Equal(t, "org-real", org.ID) + assert.Equal(t, "Real Org", org.Name) +} + func TestGetOrganizationsFiltersNameCaseInsensitive(t *testing.T) { fs := MakeMockAuthHTTPStore() httpmock.ActivateNonDefault(fs.authHTTPClient.restyClient.GetClient()) From 3ae9be1ef6b414ee0dd3e1fd19377b608d86c58d Mon Sep 17 00:00:00 2001 From: Drew Malin Date: Mon, 31 Aug 2026 13:38:28 -0700 Subject: [PATCH 2/4] minor cleanup of org resolution --- pkg/auth/auth.go | 18 +++++++----------- pkg/auth/auth_test.go | 4 ++-- pkg/cmd/register/register.go | 6 +----- pkg/cmd/register/register_test.go | 6 +++--- 4 files changed, 13 insertions(+), 21 deletions(-) diff --git a/pkg/auth/auth.go b/pkg/auth/auth.go index d385ea621..91aa41c32 100644 --- a/pkg/auth/auth.go +++ b/pkg/auth/auth.go @@ -162,26 +162,22 @@ func IsAPIKeyAuthStore(authTokensProvider APIKeyAuthStore) bool { return IsBrevAPIKey(tokens.APIKey) } -func SingleOrgForAPIKey(orgs []entity.Organization) (*entity.Organization, error) { - if len(orgs) != 1 { - return nil, breverrors.New("api key invalid") - } - return &orgs[0], nil -} - func ResolveEnvAPIKeyOrg(orgLister OrgLister) (*entity.Organization, error) { if strings.TrimSpace(os.Getenv(APIKeyEnvVar)) == "" { return nil, nil } + return ResolveAPIKeyOrganization(orgLister) +} + +func ResolveAPIKeyOrganization(orgLister OrgLister) (*entity.Organization, error) { orgs, err := orgLister.ListOrganizations() if err != nil { return nil, breverrors.WrapAndTrace(err) } - org, err := SingleOrgForAPIKey(orgs) - if err != nil { - return nil, err + if len(orgs) != 1 { + return nil, breverrors.Errorf("expected API key to resolve to exactly one organization, got %d", len(orgs)) } - return org, nil + return &orgs[0], nil } func GetAPIKeyOrgID(authTokensProvider APIKeyAuthStore) (string, error) { diff --git a/pkg/auth/auth_test.go b/pkg/auth/auth_test.go index 03664e772..f5453d7cf 100644 --- a/pkg/auth/auth_test.go +++ b/pkg/auth/auth_test.go @@ -153,7 +153,7 @@ func TestResolveEnvAPIKeyOrg_NoOrgReturnsError(t *testing.T) { s := &sideEffectingTokenStore{} _, err := ResolveEnvAPIKeyOrg(s) assert.Error(t, err) - assert.Contains(t, err.Error(), "api key invalid") + assert.Contains(t, err.Error(), "expected API key to resolve to exactly one organization, got 0") } func TestResolveEnvAPIKeyOrg_MultipleOrgsReturnsError(t *testing.T) { @@ -163,7 +163,7 @@ func TestResolveEnvAPIKeyOrg_MultipleOrgsReturnsError(t *testing.T) { }} _, err := ResolveEnvAPIKeyOrg(s) assert.Error(t, err) - assert.Contains(t, err.Error(), "api key invalid") + assert.Contains(t, err.Error(), "expected API key to resolve to exactly one organization, got 2") } func TestResolveEnvAPIKeyOrg_ListErrorPropagates(t *testing.T) { diff --git a/pkg/cmd/register/register.go b/pkg/cmd/register/register.go index c83de152b..e1a9366e4 100644 --- a/pkg/cmd/register/register.go +++ b/pkg/cmd/register/register.go @@ -391,11 +391,7 @@ func resolveOrg(s RegisterStore, orgName string) (*entity.Organization, error) { } func ResolveOrgForAPIKey(s auth.OrgLister, orgName string) (*entity.Organization, error) { - orgs, err := s.ListOrganizations() - if err != nil { - return nil, breverrors.WrapAndTrace(err) - } - org, err := auth.SingleOrgForAPIKey(orgs) + org, err := auth.ResolveAPIKeyOrganization(s) if err != nil { return nil, breverrors.WrapAndTrace(err) } diff --git a/pkg/cmd/register/register_test.go b/pkg/cmd/register/register_test.go index 0124bfb2b..092a8cd68 100644 --- a/pkg/cmd/register/register_test.go +++ b/pkg/cmd/register/register_test.go @@ -783,9 +783,9 @@ func Test_resolveOrgForAPIKey(t *testing.T) { {"single, name matches", []entity.Organization{{ID: "org_1", Name: "Alpha"}}, "Alpha", "org_1", ""}, {"single, no name", []entity.Organization{{ID: "org_1", Name: "Alpha"}}, "", "org_1", ""}, {"single, name mismatch", []entity.Organization{{ID: "org_1", Name: "Alpha"}}, "Beta", "", "does not belong to organization"}, - {"empty", nil, "", "", "api key invalid"}, - {"multiple, name matches one", []entity.Organization{{ID: "org_1", Name: "Alpha"}, {ID: "org_2", Name: "Beta"}}, "Beta", "", "api key invalid"}, - {"multiple, no name", []entity.Organization{{ID: "org_1", Name: "Alpha"}, {ID: "org_2", Name: "Beta"}}, "", "", "api key invalid"}, + {"empty", nil, "", "", "expected API key to resolve to exactly one organization, got 0"}, + {"multiple, name matches one", []entity.Organization{{ID: "org_1", Name: "Alpha"}, {ID: "org_2", Name: "Beta"}}, "Beta", "", "expected API key to resolve to exactly one organization, got 2"}, + {"multiple, no name", []entity.Organization{{ID: "org_1", Name: "Alpha"}, {ID: "org_2", Name: "Beta"}}, "", "", "expected API key to resolve to exactly one organization, got 2"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { From 24c50cea03eca47930d5d93223322ec67ec061c5 Mon Sep 17 00:00:00 2001 From: Drew Malin Date: Mon, 31 Aug 2026 13:40:18 -0700 Subject: [PATCH 3/4] restore error message --- pkg/auth/auth.go | 2 +- pkg/auth/auth_test.go | 4 ++-- pkg/cmd/register/register_test.go | 6 +++--- 3 files changed, 6 insertions(+), 6 deletions(-) diff --git a/pkg/auth/auth.go b/pkg/auth/auth.go index 91aa41c32..09f163f1d 100644 --- a/pkg/auth/auth.go +++ b/pkg/auth/auth.go @@ -175,7 +175,7 @@ func ResolveAPIKeyOrganization(orgLister OrgLister) (*entity.Organization, error return nil, breverrors.WrapAndTrace(err) } if len(orgs) != 1 { - return nil, breverrors.Errorf("expected API key to resolve to exactly one organization, got %d", len(orgs)) + return nil, breverrors.New("api key invalid") } return &orgs[0], nil } diff --git a/pkg/auth/auth_test.go b/pkg/auth/auth_test.go index f5453d7cf..03664e772 100644 --- a/pkg/auth/auth_test.go +++ b/pkg/auth/auth_test.go @@ -153,7 +153,7 @@ func TestResolveEnvAPIKeyOrg_NoOrgReturnsError(t *testing.T) { s := &sideEffectingTokenStore{} _, err := ResolveEnvAPIKeyOrg(s) assert.Error(t, err) - assert.Contains(t, err.Error(), "expected API key to resolve to exactly one organization, got 0") + assert.Contains(t, err.Error(), "api key invalid") } func TestResolveEnvAPIKeyOrg_MultipleOrgsReturnsError(t *testing.T) { @@ -163,7 +163,7 @@ func TestResolveEnvAPIKeyOrg_MultipleOrgsReturnsError(t *testing.T) { }} _, err := ResolveEnvAPIKeyOrg(s) assert.Error(t, err) - assert.Contains(t, err.Error(), "expected API key to resolve to exactly one organization, got 2") + assert.Contains(t, err.Error(), "api key invalid") } func TestResolveEnvAPIKeyOrg_ListErrorPropagates(t *testing.T) { diff --git a/pkg/cmd/register/register_test.go b/pkg/cmd/register/register_test.go index 092a8cd68..0124bfb2b 100644 --- a/pkg/cmd/register/register_test.go +++ b/pkg/cmd/register/register_test.go @@ -783,9 +783,9 @@ func Test_resolveOrgForAPIKey(t *testing.T) { {"single, name matches", []entity.Organization{{ID: "org_1", Name: "Alpha"}}, "Alpha", "org_1", ""}, {"single, no name", []entity.Organization{{ID: "org_1", Name: "Alpha"}}, "", "org_1", ""}, {"single, name mismatch", []entity.Organization{{ID: "org_1", Name: "Alpha"}}, "Beta", "", "does not belong to organization"}, - {"empty", nil, "", "", "expected API key to resolve to exactly one organization, got 0"}, - {"multiple, name matches one", []entity.Organization{{ID: "org_1", Name: "Alpha"}, {ID: "org_2", Name: "Beta"}}, "Beta", "", "expected API key to resolve to exactly one organization, got 2"}, - {"multiple, no name", []entity.Organization{{ID: "org_1", Name: "Alpha"}, {ID: "org_2", Name: "Beta"}}, "", "", "expected API key to resolve to exactly one organization, got 2"}, + {"empty", nil, "", "", "api key invalid"}, + {"multiple, name matches one", []entity.Organization{{ID: "org_1", Name: "Alpha"}, {ID: "org_2", Name: "Beta"}}, "Beta", "", "api key invalid"}, + {"multiple, no name", []entity.Organization{{ID: "org_1", Name: "Alpha"}, {ID: "org_2", Name: "Beta"}}, "", "", "api key invalid"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { From 633ac3241985845cef8c62898092d62024f36c3c Mon Sep 17 00:00:00 2001 From: Drew Malin Date: Mon, 31 Aug 2026 13:43:32 -0700 Subject: [PATCH 4/4] fix err message --- pkg/cmd/set/set.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pkg/cmd/set/set.go b/pkg/cmd/set/set.go index 0de3ac4f3..ec722fc5b 100644 --- a/pkg/cmd/set/set.go +++ b/pkg/cmd/set/set.go @@ -62,7 +62,7 @@ func set(orgName string, setStore SetStore) error { return fmt.Errorf("can not set orgs in a workspace") } if auth.IsAPIKeyAuthStore(setStore) { - return breverrors.NewValidationError("api key auth is scoped to the org saved during login; run brev login --api-key --org-id to change it") + return breverrors.NewValidationError("api key auth is scoped to the org saved during login; run brev login --api-key to change it") } orgs, err := setStore.GetOrganizations(&store.GetOrganizationsOptions{Name: orgName}) if err != nil {