From 6944cc9eaff6229da65f6a1d79e0595f1dd04e2f Mon Sep 17 00:00:00 2001 From: Pratik Patel Date: Thu, 27 Aug 2026 11:36:35 -0700 Subject: [PATCH 1/2] 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 1d3923de82124b33046a0d2b5726d3f9efa39f81 Mon Sep 17 00:00:00 2001 From: Pratik Patel Date: Wed, 26 Aug 2026 13:29:38 -0700 Subject: [PATCH 2/2] Rename enable-ssh to allow-ssh and add cert-authority support --- go.mod | 4 +- go.sum | 8 +- pkg/cmd/allowssh/allowssh.go | 213 ++++++++++++++++++ .../allowssh_test.go} | 175 ++++++++++++-- pkg/cmd/cmd.go | 6 +- pkg/cmd/deregister/deregister.go | 109 +++++++-- pkg/cmd/deregister/deregister_test.go | 125 +++++++++- pkg/cmd/disallowssh/disallowssh.go | 99 ++++++++ pkg/cmd/enablessh/enablessh.go | 153 ------------- pkg/cmd/grantssh/grantssh.go | 2 +- pkg/cmd/register/register.go | 8 +- pkg/sshcert/sshcert.go | 55 +++++ pkg/sshcert/sshcert_test.go | 131 +++++++++++ 13 files changed, 890 insertions(+), 198 deletions(-) create mode 100644 pkg/cmd/allowssh/allowssh.go rename pkg/cmd/{enablessh/enablessh_test.go => allowssh/allowssh_test.go} (60%) create mode 100644 pkg/cmd/disallowssh/disallowssh.go delete mode 100644 pkg/cmd/enablessh/enablessh.go diff --git a/go.mod b/go.mod index 5855b1103..83b067310 100644 --- a/go.mod +++ b/go.mod @@ -3,8 +3,8 @@ module github.com/brevdev/brev-cli go 1.25.0 require ( - buf.build/gen/go/brevdev/devplane/connectrpc/go v1.20.0-20260820222245-1cfc91443320.1 - buf.build/gen/go/brevdev/devplane/protocolbuffers/go v1.36.12-20260820222245-1cfc91443320.1 + buf.build/gen/go/brevdev/devplane/connectrpc/go v1.20.0-20260827214152-35c65570f2a0.1 + buf.build/gen/go/brevdev/devplane/protocolbuffers/go v1.36.12-20260827214152-35c65570f2a0.1 connectrpc.com/connect v1.20.0 github.com/NVIDIA/go-nvml v0.13.0-1 github.com/alessio/shellescape v1.4.1 diff --git a/go.sum b/go.sum index 4af12529f..f8eb670c1 100644 --- a/go.sum +++ b/go.sum @@ -1,7 +1,7 @@ -buf.build/gen/go/brevdev/devplane/connectrpc/go v1.20.0-20260820222245-1cfc91443320.1 h1:PKIsaGilewnQUSHNUn+Ir4sagWne713vJS3Ys7h9vAY= -buf.build/gen/go/brevdev/devplane/connectrpc/go v1.20.0-20260820222245-1cfc91443320.1/go.mod h1:r4xfuOy9bpAXm13ugDRO+JNmFVlXecGRuKtn1X7os/k= -buf.build/gen/go/brevdev/devplane/protocolbuffers/go v1.36.12-20260820222245-1cfc91443320.1 h1:gmAgE9NC+BAovZIs9CNmjgExqM+Gox8AZ6ud3eVMxfA= -buf.build/gen/go/brevdev/devplane/protocolbuffers/go v1.36.12-20260820222245-1cfc91443320.1/go.mod h1:N18pnR0HL6srurI7G19FpSEki71wA1u4e2c5zbfeTV8= +buf.build/gen/go/brevdev/devplane/connectrpc/go v1.20.0-20260827214152-35c65570f2a0.1 h1:xzM4gdexDMGgTwdgrlUFHgedW+CSbdXaQzOimV5PbPU= +buf.build/gen/go/brevdev/devplane/connectrpc/go v1.20.0-20260827214152-35c65570f2a0.1/go.mod h1:qMKDH/phd8XN/OWkJlSVhHJ/8P2w1dfE5zUQprRRp8c= +buf.build/gen/go/brevdev/devplane/protocolbuffers/go v1.36.12-20260827214152-35c65570f2a0.1 h1:17qqLaEUl7Biv3eyy9owm5yrS/Zi9Zyxc01XSZIzVbU= +buf.build/gen/go/brevdev/devplane/protocolbuffers/go v1.36.12-20260827214152-35c65570f2a0.1/go.mod h1:N18pnR0HL6srurI7G19FpSEki71wA1u4e2c5zbfeTV8= buf.build/gen/go/brevdev/protoc-gen-gotag/protocolbuffers/go v1.36.12-20220906235457-8b4922735da5.1 h1:Qk/4GJyWVWvWsfEFeX4T+k7KouZdRUxxUnIUwJ3hmZg= buf.build/gen/go/brevdev/protoc-gen-gotag/protocolbuffers/go v1.36.12-20220906235457-8b4922735da5.1/go.mod h1:SacJAYqnICCQAsBA46cSA/hxhqhxYkiYzseucf6/fhQ= cloud.google.com/go v0.26.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMTw= diff --git a/pkg/cmd/allowssh/allowssh.go b/pkg/cmd/allowssh/allowssh.go new file mode 100644 index 000000000..60e997b6f --- /dev/null +++ b/pkg/cmd/allowssh/allowssh.go @@ -0,0 +1,213 @@ +// Package allowssh implements brev allow-ssh. +package allowssh + +import ( + "context" + "fmt" + "os" + "os/exec" + "os/user" + "path/filepath" + "strings" + + nodev1 "buf.build/gen/go/brevdev/devplane/protocolbuffers/go/devplaneapi/v1" + "connectrpc.com/connect" + + "github.com/brevdev/brev-cli/pkg/cmd/register" + "github.com/brevdev/brev-cli/pkg/config" + "github.com/brevdev/brev-cli/pkg/entity" + "github.com/brevdev/brev-cli/pkg/externalnode" + "github.com/brevdev/brev-cli/pkg/sshcert" + "github.com/brevdev/brev-cli/pkg/terminal" + + "github.com/spf13/cobra" +) + +type AllowSSHStore interface { + GetCurrentUser() (*entity.User, error) + GetAccessToken() (string, error) +} + +type allowSSHDeps struct { + platform externalnode.PlatformChecker + nodeClients externalnode.NodeClientFactory + registrationStore register.RegistrationStore + prompter terminal.Selector +} + +func defaultAllowSSHDeps() allowSSHDeps { + return allowSSHDeps{ + platform: register.LinuxPlatform{}, + nodeClients: register.DefaultNodeClientFactory{}, + registrationStore: register.NewFileRegistrationStore(), + prompter: register.TerminalPrompter{}, + } +} + +func NewCmdAllowSSH(t *terminal.Terminal, store AllowSSHStore) *cobra.Command { + cmd := &cobra.Command{ + Annotations: map[string]string{"configuration": ""}, + Use: "allow-ssh", + DisableFlagsInUseLine: true, + Short: "Trust the Brev certificate authority on this device for SSH", + Long: "Writes the Brev certificate authority to authorized_keys, allowing this device to be an SSH target for the current Linux user. Users are granted access with 'brev grant-ssh'.", + Example: " brev allow-ssh", + RunE: func(cmd *cobra.Command, args []string) error { + return runAllowSSH(cmd.Context(), t, store, defaultAllowSSHDeps()) + }, + } + + return cmd +} + +func runAllowSSH(ctx context.Context, t *terminal.Terminal, s AllowSSHStore, deps allowSSHDeps) error { + if !deps.platform.IsCompatible() { + return fmt.Errorf("brev allow-ssh is only supported on Linux") + } + + reg, err := deps.registrationStore.Load() + if err != nil { + return fmt.Errorf("failed to read registration file: %w", err) + } + + brevUser, err := s.GetCurrentUser() + if err != nil { + return fmt.Errorf("failed to get current user: %w", err) + } + + return allowSSH(ctx, t, deps, s, reg, brevUser) +} + +func allowSSH( + ctx context.Context, + t *terminal.Terminal, + deps allowSSHDeps, + tokenProvider externalnode.TokenProvider, + reg *register.DeviceRegistration, + brevUser *entity.User, +) error { + linuxUser, err := user.Current() + if err != nil { + return fmt.Errorf("failed to determine current Linux user: %w", err) + } + linuxUsername := linuxUser.Username + + checkSSHDaemon(t) + + t.Vprint("") + t.Vprint(t.Green("Allowing SSH on this device")) + t.Vprint("") + t.Vprintf(" Node: %s (%s)\n", reg.DisplayName, reg.ExternalNodeID) + t.Vprintf(" Linux user: %s\n", linuxUsername) + t.Vprint("") + + node, err := fetchRegisteredNode(ctx, deps, tokenProvider, reg) + if err != nil { + return fmt.Errorf("allow SSH failed: %w", err) + } + + if node.GetLabels()[sshcert.LabelKeySSHProvider] != sshcert.SSHProviderCertAuth { + return legacyEnableSSH(ctx, t, deps, tokenProvider, reg, brevUser, node, linuxUsername) + } + + caPublicKey := node.GetCertificateAuthority() + + if err := installCertAuthority(linuxUser, caPublicKey, reg.ExternalNodeID, linuxUsername); err != nil { + return fmt.Errorf("allow SSH failed: %w", err) + } + t.Vprint(t.Green(" Certificate authority written to authorized_keys.")) + + t.Vprint("") + t.Vprint(t.Green("SSH allowed on this device. No one has SSH access yet — grant it with: brev grant-ssh")) + return nil +} + +func legacyEnableSSH( + ctx context.Context, + t *terminal.Terminal, + deps allowSSHDeps, + tokenProvider externalnode.TokenProvider, + reg *register.DeviceRegistration, + brevUser *entity.User, + node *nodev1.ExternalNode, + linuxUsername string, +) error { + brevPortID, err := register.ResolveSSHAccessPort(ctx, t, deps.prompter, deps.nodeClients, tokenProvider, reg, node) + if err != nil { + return fmt.Errorf("allow SSH failed: %w", err) + } + + if err := register.SetupAndRegisterNodeSSHAccess(ctx, t, deps.nodeClients, tokenProvider, reg, brevUser, linuxUsername, brevPortID); err != nil { + return fmt.Errorf("allow SSH failed: %w", err) + } + + t.Vprint("") + t.Vprint(t.Green(fmt.Sprintf("SSH access enabled. You can now SSH to this device via: brev shell %s", reg.DisplayName))) + return nil +} + +func installCertAuthority(osUser *user.User, caPublicKey, nodeID, linuxUser string) error { + if caPublicKey == "" { + return fmt.Errorf("certificate authority public key is required") + } + + principal := fmt.Sprintf("brev:v1:vm:%s:login:%s", nodeID, linuxUser) + entry := fmt.Sprintf("cert-authority,principals=\"%s\" %s", principal, strings.TrimSpace(caPublicKey)) + + sshDir := filepath.Join(osUser.HomeDir, ".ssh") + if err := os.MkdirAll(sshDir, 0o700); err != nil { + return fmt.Errorf("creating .ssh directory: %w", err) + } + + authKeysPath := filepath.Join(sshDir, "authorized_keys") + + existing, err := os.ReadFile(authKeysPath) // #nosec G304 + if err != nil && !os.IsNotExist(err) { + return fmt.Errorf("reading authorized_keys: %w", err) + } + + // skip if the entry already exists. + for line := range strings.SplitSeq(string(existing), "\n") { + if strings.TrimSpace(line) == entry { + return nil + } + } + + content := string(existing) + if content != "" && !strings.HasSuffix(content, "\n") { + content += "\n" + } + content += entry + "\n" + + if err := os.WriteFile(authKeysPath, []byte(content), 0o600); err != nil { + return fmt.Errorf("writing authorized_keys: %w", err) + } + + return nil +} + +func fetchRegisteredNode( + ctx context.Context, + deps allowSSHDeps, + tokenProvider externalnode.TokenProvider, + reg *register.DeviceRegistration, +) (*nodev1.ExternalNode, error) { + client := deps.nodeClients.NewNodeClient(tokenProvider, config.GlobalConfig.GetBrevPublicAPIURL()) + resp, err := client.GetNode(ctx, connect.NewRequest(&nodev1.GetNodeRequest{ + ExternalNodeId: reg.ExternalNodeID, + })) + if err != nil { + return nil, fmt.Errorf("error retrieving node: %w", err) + } + return resp.Msg.GetExternalNode(), nil +} + +func checkSSHDaemon(t *terminal.Terminal) { + for _, svc := range []string{"ssh", "sshd"} { + out, err := exec.Command("systemctl", "is-active", svc).Output() //nolint:gosec // fixed service names + if err == nil && len(out) > 0 && string(out[:len(out)-1]) == "active" { + return + } + } + t.Vprintf(" %s\n", t.Yellow("Warning: SSH daemon does not appear to be running. SSH access may not work until sshd is started.")) +} diff --git a/pkg/cmd/enablessh/enablessh_test.go b/pkg/cmd/allowssh/allowssh_test.go similarity index 60% rename from pkg/cmd/enablessh/enablessh_test.go rename to pkg/cmd/allowssh/allowssh_test.go index 7df94144d..fb0977f22 100644 --- a/pkg/cmd/enablessh/enablessh_test.go +++ b/pkg/cmd/allowssh/allowssh_test.go @@ -1,7 +1,8 @@ -package enablessh +package allowssh import ( "context" + "fmt" "net/http/httptest" "os" "os/user" @@ -14,16 +15,16 @@ import ( "connectrpc.com/connect" "github.com/brevdev/brev-cli/pkg/cmd/register" + "github.com/brevdev/brev-cli/pkg/entity" "github.com/brevdev/brev-cli/pkg/externalnode" + "github.com/brevdev/brev-cli/pkg/terminal" ) -// tempUser returns a *user.User whose HomeDir points to a temporary directory. func tempUser(t *testing.T) *user.User { t.Helper() return &user.User{HomeDir: t.TempDir()} } -// readAuthorizedKeys is a test helper that reads ~/.ssh/authorized_keys. func readAuthorizedKeys(t *testing.T, u *user.User) string { t.Helper() data, err := os.ReadFile(filepath.Join(u.HomeDir, ".ssh", "authorized_keys")) @@ -224,17 +225,51 @@ func (m mockNodeClientFactory) NewNodeClient(provider externalnode.TokenProvider return register.NewNodeServiceClient(provider, m.serverURL) } -type mockEnableSSHStore struct { +type mockAllowSSHStore struct { token string } -func (m *mockEnableSSHStore) GetCurrentUser() (interface{}, error) { return nil, nil } -func (m *mockEnableSSHStore) GetAccessToken() (string, error) { return m.token, nil } +func (m *mockAllowSSHStore) GetCurrentUser() (*entity.User, error) { return &entity.User{}, nil } +func (m *mockAllowSSHStore) GetAccessToken() (string, error) { return m.token, nil } + +// mockSelector implements terminal.Selector, returning the first item. +type mockSelector struct{ choice string } + +func (m mockSelector) Select(_ string, items []string) string { + if m.choice != "" { + for _, s := range items { + if s == m.choice { + return s + } + } + } + if len(items) > 0 { + return items[0] + } + return "" +} -// fakeNodeService implements the server side of ExternalNodeService for testing. type fakeNodeService struct { nodev1connect.UnimplementedExternalNodeServiceHandler - getNodeFn func(*nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) + getNodeFn func(*nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) + grantCalls int + openCalls int +} + +func (f *fakeNodeService) GrantNodeSSHAccess(_ context.Context, _ *connect.Request[nodev1.GrantNodeSSHAccessRequest]) (*connect.Response[nodev1.GrantNodeSSHAccessResponse], error) { + f.grantCalls++ + return connect.NewResponse(&nodev1.GrantNodeSSHAccessResponse{}), nil +} + +func (f *fakeNodeService) OpenPort(_ context.Context, req *connect.Request[nodev1.OpenPortRequest]) (*connect.Response[nodev1.OpenPortResponse], error) { + f.openCalls++ + return connect.NewResponse(&nodev1.OpenPortResponse{ + Port: &nodev1.Port{ + PortId: fmt.Sprintf("port_%d", req.Msg.GetPortNumber()), + Protocol: req.Msg.GetProtocol(), + PortNumber: req.Msg.GetPortNumber(), + }, + }), nil } func (f *fakeNodeService) GetNode(_ context.Context, req *connect.Request[nodev1.GetNodeRequest]) (*connect.Response[nodev1.GetNodeResponse], error) { @@ -245,14 +280,15 @@ func (f *fakeNodeService) GetNode(_ context.Context, req *connect.Request[nodev1 return connect.NewResponse(resp), nil } -func startFakeServer(t *testing.T, svc *fakeNodeService) (enableSSHDeps, *httptest.Server) { +func startFakeServer(t *testing.T, svc *fakeNodeService) allowSSHDeps { t.Helper() _, handler := nodev1connect.NewExternalNodeServiceHandler(svc) server := httptest.NewServer(handler) t.Cleanup(server.Close) - return enableSSHDeps{ + return allowSSHDeps{ nodeClients: mockNodeClientFactory{serverURL: server.URL}, - }, server + prompter: mockSelector{}, + } } func Test_fetchRegisteredNode(t *testing.T) { @@ -267,8 +303,8 @@ func Test_fetchRegisteredNode(t *testing.T) { }}, nil }, } - deps, _ := startFakeServer(t, svc) - store := &mockEnableSSHStore{token: "tok"} + deps := startFakeServer(t, svc) + store := &mockAllowSSHStore{token: "tok"} reg := ®ister.DeviceRegistration{ExternalNodeID: "unode_abc", OrgID: "org_1"} node, err := fetchRegisteredNode(context.Background(), deps, store, reg) @@ -279,3 +315,116 @@ func Test_fetchRegisteredNode(t *testing.T) { t.Fatalf("unexpected node: %+v", node) } } + +// --- installCertAuthority --- + +func Test_installCertAuthority(t *testing.T) { + const ( + caKey = "ssh-ed25519 AAAAC3Nz dummyCA" + node = "unode_abc" + luser = "ubuntu" + ) + + t.Run("WritesLine", func(t *testing.T) { + u := tempUser(t) + if err := installCertAuthority(u, caKey, node, luser); err != nil { + t.Fatalf("installCertAuthority: %v", err) + } + want := `cert-authority,principals="brev:v1:vm:unode_abc:login:ubuntu" ssh-ed25519 AAAAC3Nz dummyCA` + if result := readAuthorizedKeys(t, u); !strings.Contains(result, want) { + t.Errorf("expected cert-authority line not found:\n%s", result) + } + }) + + t.Run("Idempotent", func(t *testing.T) { + u := tempUser(t) + for i := 0; i < 2; i++ { + if err := installCertAuthority(u, caKey, node, luser); err != nil { + t.Fatalf("installCertAuthority #%d: %v", i+1, err) + } + } + result := readAuthorizedKeys(t, u) + if n := strings.Count(result, "cert-authority"); n != 1 { + t.Errorf("expected 1 cert-authority line, got %d:\n%s", n, result) + } + }) + + t.Run("PreservesExistingKeys", func(t *testing.T) { + u := tempUser(t) + sshDir := filepath.Join(u.HomeDir, ".ssh") + if err := os.MkdirAll(sshDir, 0o700); err != nil { + t.Fatal(err) + } + original := "ssh-rsa EXISTING user@host\n" + if err := os.WriteFile(filepath.Join(sshDir, "authorized_keys"), []byte(original), 0o600); err != nil { + t.Fatal(err) + } + + if err := installCertAuthority(u, caKey, node, luser); err != nil { + t.Fatalf("installCertAuthority: %v", err) + } + + result := readAuthorizedKeys(t, u) + if !strings.Contains(result, "ssh-rsa EXISTING user@host") { + t.Errorf("existing key was removed:\n%s", result) + } + if !strings.Contains(result, "cert-authority") { + t.Errorf("cert-authority line not written:\n%s", result) + } + }) + + t.Run("EmptyKeyErrors", func(t *testing.T) { + if err := installCertAuthority(tempUser(t), "", node, luser); err == nil { + t.Error("expected error for empty CA key") + } + }) +} + +func Test_allowSSH_LegacyNodeFallsBackToKeys(t *testing.T) { + svc := &fakeNodeService{ + getNodeFn: func(_ *nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) { + return &nodev1.GetNodeResponse{ + ExternalNode: &nodev1.ExternalNode{ + ExternalNodeId: "unode_legacy", + // No sshprovider label — legacy node. + Labels: map[string]string{}, + Ports: []*nodev1.Port{{ + PortId: "port_ssh", + Protocol: nodev1.PortProtocol_PORT_PROTOCOL_TCP, + PortNumber: 22, + }}, + }, + }, nil + }, + } + deps := startFakeServer(t, svc) + + reg := ®ister.DeviceRegistration{ + DisplayName: "legacy-node", + ExternalNodeID: "unode_legacy", + OrgID: "org_1", + } + + term := terminal.New() + if err := allowSSH(context.Background(), term, deps, &mockAllowSSHStore{}, reg, &entity.User{ID: "user_1"}); err != nil { + t.Fatalf("allowSSH failed: %v", err) + } + + // Legacy flow must grant SSH access (reflexive grant). + if svc.grantCalls == 0 { + t.Error("expected GrantNodeSSHAccess to be called for legacy node") + } + + // No cert-authority line may be written for a legacy node. + realUser, err := user.Current() + if err != nil { + t.Fatalf("user.Current failed: %v", err) + } + authKeysPath := filepath.Join(realUser.HomeDir, ".ssh", "authorized_keys") + data, readErr := os.ReadFile(authKeysPath) // #nosec G304 + if readErr == nil { + if strings.Contains(string(data), "brev:v1:vm:unode_legacy") { + t.Errorf("legacy node must not write a cert-authority line:\n%s", string(data)) + } + } +} diff --git a/pkg/cmd/cmd.go b/pkg/cmd/cmd.go index a845788cd..09ebbff66 100644 --- a/pkg/cmd/cmd.go +++ b/pkg/cmd/cmd.go @@ -8,6 +8,7 @@ import ( "github.com/brevdev/brev-cli/pkg/analytics" "github.com/brevdev/brev-cli/pkg/auth" "github.com/brevdev/brev-cli/pkg/cmd/agentskill" + "github.com/brevdev/brev-cli/pkg/cmd/allowssh" analyticscmd "github.com/brevdev/brev-cli/pkg/cmd/analytics" "github.com/brevdev/brev-cli/pkg/cmd/background" "github.com/brevdev/brev-cli/pkg/cmd/clipboard" @@ -16,7 +17,7 @@ import ( "github.com/brevdev/brev-cli/pkg/cmd/copy" "github.com/brevdev/brev-cli/pkg/cmd/delete" "github.com/brevdev/brev-cli/pkg/cmd/deregister" - "github.com/brevdev/brev-cli/pkg/cmd/enablessh" + "github.com/brevdev/brev-cli/pkg/cmd/disallowssh" "github.com/brevdev/brev-cli/pkg/cmd/envvars" "github.com/brevdev/brev-cli/pkg/cmd/exec" "github.com/brevdev/brev-cli/pkg/cmd/feedback" @@ -333,7 +334,8 @@ func createCmdTree(cmd *cobra.Command, t *terminal.Terminal, loginCmdStore *stor cmd.AddCommand(register.NewCmdRegister(t, externalNodeCmdStore)) cmd.AddCommand(deregister.NewCmdDeregister(t, externalNodeCmdStore)) cmd.AddCommand(upgrade.NewCmdUpgrade(t, noLoginCmdStore)) - cmd.AddCommand(enablessh.NewCmdEnableSSH(t, externalNodeCmdStore)) + cmd.AddCommand(allowssh.NewCmdAllowSSH(t, externalNodeCmdStore)) + cmd.AddCommand(disallowssh.NewCmdDisallowSSH(t, externalNodeCmdStore)) cmd.AddCommand(grantssh.NewCmdGrantSSH(t, externalNodeCmdStore)) cmd.AddCommand(revokessh.NewCmdRevokeSSH(t, externalNodeCmdStore)) cmd.AddCommand(runtasks.NewCmdRunTasks(t, noLoginCmdStore)) diff --git a/pkg/cmd/deregister/deregister.go b/pkg/cmd/deregister/deregister.go index eb0d16e13..4fb335d45 100644 --- a/pkg/cmd/deregister/deregister.go +++ b/pkg/cmd/deregister/deregister.go @@ -15,6 +15,7 @@ import ( "github.com/brevdev/brev-cli/pkg/config" "github.com/brevdev/brev-cli/pkg/entity" "github.com/brevdev/brev-cli/pkg/externalnode" + "github.com/brevdev/brev-cli/pkg/sshcert" "github.com/brevdev/brev-cli/pkg/sudo" "github.com/brevdev/brev-cli/pkg/terminal" @@ -26,18 +27,27 @@ type DeregisterStore interface { GetAccessToken() (string, error) } -type SSHKeyRemover interface { +type CertAuthorityRemover interface { + RemoveCertAuthority(u *user.User, nodeID, linuxUser string) (bool, error) +} + +// LegacySSHKeyRemover removes Brev-managed per-user SSH keys (legacy nodes). +type LegacySSHKeyRemover interface { RemoveBrevKeys(u *user.User) ([]string, error) } -type brevSSHKeyRemover struct{} +type brevCertAuthorityRemover struct{} + +func (brevCertAuthorityRemover) RemoveCertAuthority(u *user.User, nodeID, linuxUser string) (bool, error) { + removed, err := sshcert.RemoveCertAuthorityLine(u.HomeDir, nodeID, linuxUser) + return removed, breverrors.WrapAndTrace(err) +} -func (brevSSHKeyRemover) RemoveBrevKeys(u *user.User) ([]string, error) { +type legacyKeyRemover struct{} + +func (legacyKeyRemover) RemoveBrevKeys(u *user.User) ([]string, error) { removed, err := register.RemoveBrevAuthorizedKeys(u) - if err != nil { - return nil, fmt.Errorf("removing brev authorized keys: %w", err) - } - return removed, nil + return removed, breverrors.WrapAndTrace(err) } // deregisterDeps bundles the side-effecting dependencies of runDeregister so @@ -50,7 +60,8 @@ type deregisterDeps struct { netbird register.NetBirdManager nodeClients externalnode.NodeClientFactory registrationStore register.RegistrationStore - sshKeys SSHKeyRemover + sshKeys CertAuthorityRemover + legacyKeys LegacySSHKeyRemover } func defaultDeregisterDeps() deregisterDeps { @@ -60,9 +71,9 @@ func defaultDeregisterDeps() deregisterDeps { confirmer: register.TerminalPrompter{}, gater: sudo.Default, netbird: register.Netbird{}, - nodeClients: register.DefaultNodeClientFactory{}, + sshKeys: brevCertAuthorityRemover{}, + legacyKeys: legacyKeyRemover{}, registrationStore: register.NewFileRegistrationStore(), - sshKeys: brevSSHKeyRemover{}, } } @@ -197,7 +208,7 @@ func runDeregister(ctx context.Context, t *terminal.Terminal, s DeregisterStore, t.Vprint("") t.Vprint(t.Yellow(" This will:")) t.Vprint(" 1. Remove this node from Brev") - t.Vprint(" 2. Remove Brev SSH keys from this machine (if any)") + t.Vprint(" 2. Remove any SSH data associated with this node") t.Vprint(" 3. Uninstall the Brev tunnel") t.Vprint(" 4. Delete local registration data") t.Vprint("") @@ -219,21 +230,21 @@ func runDeregister(ctx context.Context, t *terminal.Terminal, s DeregisterStore, } t.Vprint("") - t.Vprint(t.Yellow("[Step 2/4] Removing Brev SSH keys...")) + t.Vprint(t.Yellow("[Step 2/4] Removing any SSH data associated with this node...")) if osUser == nil { t.Vprintf(" %s\n", t.Yellow("Skipped: could not determine current user")) } else { - removed, kerr := deps.sshKeys.RemoveBrevKeys(osUser) + linuxUsername := osUser.Username + certAuth, nerr := nodeUsesCertAuth(ctx, s, deps, reg) switch { - case kerr != nil: - t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: failed to remove Brev SSH keys: %v", kerr))) - case len(removed) > 0: - t.Vprintf("%s Brev SSH keys removed from authorized_keys:\n", t.Green(" ✓")) - for _, key := range removed { - t.Vprintf(" - %s\n", key) - } + case nerr != nil: + t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: could not determine node SSH mode; cleaning up both: %v", nerr))) + removeCertAuthorityStep(t, deps, osUser, reg.ExternalNodeID, linuxUsername) + removeLegacyKeysStep(t, deps, osUser) + case certAuth: + removeCertAuthorityStep(t, deps, osUser, reg.ExternalNodeID, linuxUsername) default: - t.Vprint(" No Brev SSH keys found in authorized_keys.") + removeLegacyKeysStep(t, deps, osUser) } } t.Vprint("") @@ -260,3 +271,59 @@ func runDeregister(ctx context.Context, t *terminal.Terminal, s DeregisterStore, return nil } + +// nodeUsesCertAuth reports whether the node has the sshprovider=certauth label. +// On lookup failure it falls back to checking the local authorized_keys for a +// Brev cert-authority line (so certauth nodes still get the right cleanup). +func nodeUsesCertAuth(ctx context.Context, s externalnode.TokenProvider, deps deregisterDeps, reg *register.DeviceRegistration) (bool, error) { + client := deps.nodeClients.NewNodeClient(s, config.GlobalConfig.GetBrevPublicAPIURL()) + resp, err := client.GetNode(ctx, connect.NewRequest(&nodev1.GetNodeRequest{ + ExternalNodeId: reg.ExternalNodeID, + })) + if err != nil { + // Node may already be gone (removed in step 1) or unreachable: infer + // from local state so cleanup still picks the right mode. + if hasLocalCertAuthority(reg.ExternalNodeID) { + return true, breverrors.WrapAndTrace(err) + } + return false, breverrors.WrapAndTrace(err) + } + return resp.Msg.GetExternalNode().GetLabels()[sshcert.LabelKeySSHProvider] == sshcert.SSHProviderCertAuth, nil +} + +// hasLocalCertAuthority checks ~/.ssh/authorized_keys for a Brev cert-authority +// line belonging to the given node. +func hasLocalCertAuthority(nodeID string) bool { + osUser, err := user.Current() + if err != nil { + return false + } + return sshcert.HasCertAuthorityLine(osUser.HomeDir, nodeID) +} + +func removeCertAuthorityStep(t *terminal.Terminal, deps deregisterDeps, osUser *user.User, nodeID, linuxUser string) { + removed, cerr := deps.sshKeys.RemoveCertAuthority(osUser, nodeID, linuxUser) + switch { + case cerr != nil: + t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: failed to remove cert-authority: %v", cerr))) + case removed: + t.Vprintf("%s Certificate authority removed from authorized_keys.\n", t.Green(" ✓")) + default: + t.Vprint(" No certificate authority line found in authorized_keys.") + } +} + +func removeLegacyKeysStep(t *terminal.Terminal, deps deregisterDeps, osUser *user.User) { + removed, kerr := deps.legacyKeys.RemoveBrevKeys(osUser) + switch { + case kerr != nil: + t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: failed to remove Brev SSH keys: %v", kerr))) + case len(removed) > 0: + t.Vprintf("%s Brev SSH keys removed from authorized_keys:\n", t.Green(" ✓")) + for _, key := range removed { + t.Vprintf(" - %s\n", key) + } + default: + t.Vprint(" No Brev SSH keys found in authorized_keys.") + } +} diff --git a/pkg/cmd/deregister/deregister_test.go b/pkg/cmd/deregister/deregister_test.go index 781d8de48..7d3a30794 100644 --- a/pkg/cmd/deregister/deregister_test.go +++ b/pkg/cmd/deregister/deregister_test.go @@ -39,6 +39,24 @@ type fakeNodeService struct { nodev1connect.UnimplementedExternalNodeServiceHandler removeNodeFn func(*nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) listNodesFn func(*nodev1.ListNodesRequest) (*nodev1.ListNodesResponse, error) + getNodeFn func(*nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) +} + +func (f *fakeNodeService) GetNode(_ context.Context, req *connect.Request[nodev1.GetNodeRequest]) (*connect.Response[nodev1.GetNodeResponse], error) { + if f.getNodeFn == nil { + // Default: certauth node (matches registration on this branch). + return connect.NewResponse(&nodev1.GetNodeResponse{ + ExternalNode: &nodev1.ExternalNode{ + ExternalNodeId: req.Msg.GetExternalNodeId(), + Labels: map[string]string{"sshprovider": "certauth"}, + }, + }), nil + } + resp, err := f.getNodeFn(req.Msg) + if err != nil { + return nil, err + } + return connect.NewResponse(resp), nil } func (f *fakeNodeService) RemoveNode(_ context.Context, req *connect.Request[nodev1.RemoveNodeRequest]) (*connect.Response[nodev1.RemoveNodeResponse], error) { @@ -123,12 +141,23 @@ func (m mockNodeClientFactory) NewNodeClient(provider externalnode.TokenProvider } type mockSSHKeyRemover struct { + called bool + err error + removed bool +} + +func (m *mockSSHKeyRemover) RemoveCertAuthority(_ *user.User, _, _ string) (bool, error) { + m.called = true + return m.removed, m.err +} + +type mockLegacyKeyRemover struct { called bool err error removed []string } -func (m *mockSSHKeyRemover) RemoveBrevKeys(_ *user.User) ([]string, error) { +func (m *mockLegacyKeyRemover) RemoveBrevKeys(_ *user.User) ([]string, error) { m.called = true return m.removed, m.err } @@ -179,6 +208,7 @@ func testDeregisterDeps(t *testing.T, svc *fakeNodeService, regStore register.Re nodeClients: mockNodeClientFactory{serverURL: server.URL}, registrationStore: regStore, sshKeys: &mockSSHKeyRemover{}, + legacyKeys: &mockLegacyKeyRemover{}, }, server } @@ -484,3 +514,96 @@ func Test_runDeregister_RemoveBrevKeysHandling(t *testing.T) { }) } } + +func Test_runDeregister_LegacyNodeRemovesKeys(t *testing.T) { + regStore := &mockRegistrationStore{reg: registeredReg()} + svc := &fakeNodeService{ + removeNodeFn: func(_ *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { + return &nodev1.RemoveNodeResponse{}, nil + }, + getNodeFn: func(req *nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) { + return &nodev1.GetNodeResponse{ + ExternalNode: &nodev1.ExternalNode{ + ExternalNodeId: req.GetExternalNodeId(), + // No sshprovider label — legacy node. + Labels: map[string]string{}, + }, + }, nil + }, + } + + certMock := &mockSSHKeyRemover{} + legacyMock := &mockLegacyKeyRemover{removed: []string{"ssh-rsa OLD user@host"}} + + err := runDeregisterCase(t, regStore, svc, func(d *deregisterDeps) { + d.sshKeys = certMock + d.legacyKeys = legacyMock + }) + if err != nil { + t.Fatalf("runDeregister failed: %v", err) + } + + if !legacyMock.called { + t.Error("expected RemoveBrevKeys to be called for legacy node") + } + if certMock.called { + t.Error("expected RemoveCertAuthority NOT to be called for legacy node") + } +} + +func Test_runDeregister_CertAuthNodeRemovesCertAuthority(t *testing.T) { + regStore := &mockRegistrationStore{reg: registeredReg()} + svc := &fakeNodeService{ + removeNodeFn: func(_ *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { + return &nodev1.RemoveNodeResponse{}, nil + }, + } // default getNodeFn returns a certauth node + + certMock := &mockSSHKeyRemover{removed: true} + legacyMock := &mockLegacyKeyRemover{} + + err := runDeregisterCase(t, regStore, svc, func(d *deregisterDeps) { + d.sshKeys = certMock + d.legacyKeys = legacyMock + }) + if err != nil { + t.Fatalf("runDeregister failed: %v", err) + } + + if !certMock.called { + t.Error("expected RemoveCertAuthority to be called for certauth node") + } + if legacyMock.called { + t.Error("expected RemoveBrevKeys NOT to be called for certauth node") + } +} + +func Test_runDeregister_NodeLookupFailure_CleansBoth(t *testing.T) { + regStore := &mockRegistrationStore{reg: registeredReg()} + svc := &fakeNodeService{ + removeNodeFn: func(_ *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { + return &nodev1.RemoveNodeResponse{}, nil + }, + getNodeFn: func(_ *nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) { + return nil, connect.NewError(connect.CodeInternal, fmt.Errorf("backend down")) + }, + } + + certMock := &mockSSHKeyRemover{removed: true} + legacyMock := &mockLegacyKeyRemover{removed: []string{"ssh-rsa OLD"}} + + err := runDeregisterCase(t, regStore, svc, func(d *deregisterDeps) { + d.sshKeys = certMock + d.legacyKeys = legacyMock + }) + if err != nil { + t.Fatalf("runDeregister failed: %v", err) + } + + if !certMock.called { + t.Error("expected RemoveCertAuthority to be called on lookup failure") + } + if !legacyMock.called { + t.Error("expected RemoveBrevKeys to be called on lookup failure") + } +} diff --git a/pkg/cmd/disallowssh/disallowssh.go b/pkg/cmd/disallowssh/disallowssh.go new file mode 100644 index 000000000..fca6d0e0a --- /dev/null +++ b/pkg/cmd/disallowssh/disallowssh.go @@ -0,0 +1,99 @@ +package disallowssh + +import ( + "context" + "fmt" + "os/user" + + "github.com/brevdev/brev-cli/pkg/cmd/register" + "github.com/brevdev/brev-cli/pkg/entity" + "github.com/brevdev/brev-cli/pkg/externalnode" + "github.com/brevdev/brev-cli/pkg/sshcert" + "github.com/brevdev/brev-cli/pkg/terminal" + + "github.com/spf13/cobra" +) + +type DisallowSSHStore interface { + GetCurrentUser() (*entity.User, error) + GetAccessToken() (string, error) +} + +type disallowSSHDeps struct { + platform externalnode.PlatformChecker + registrationStore register.RegistrationStore +} + +func defaultDisallowSSHDeps() disallowSSHDeps { + return disallowSSHDeps{ + platform: register.LinuxPlatform{}, + registrationStore: register.NewFileRegistrationStore(), + } +} + +func NewCmdDisallowSSH(t *terminal.Terminal, store DisallowSSHStore) *cobra.Command { + cmd := &cobra.Command{ + Annotations: map[string]string{"configuration": ""}, + Use: "disallow-ssh", + DisableFlagsInUseLine: true, + Short: "Remove Brev SSH access data from this device", + Long: "Removes the Brev certificate authority line and any Brev-managed SSH keys from authorized_keys, revoking SSH access for all users. The node remains registered.", + Example: " brev disallow-ssh", + RunE: func(cmd *cobra.Command, args []string) error { + return runDisallowSSH(cmd.Context(), t, store, defaultDisallowSSHDeps()) + }, + } + + return cmd +} + +func runDisallowSSH(_ context.Context, t *terminal.Terminal, _ DisallowSSHStore, deps disallowSSHDeps) error { + if !deps.platform.IsCompatible() { + return fmt.Errorf("brev disallow-ssh is only supported on Linux") + } + + reg, err := deps.registrationStore.Load() + if err != nil { + return fmt.Errorf("failed to read registration file: %w", err) + } + + linuxUser, err := user.Current() + if err != nil { + return fmt.Errorf("failed to determine current Linux user: %w", err) + } + + t.Vprint("") + t.Vprint(t.Green("Removing SSH certificate authority from this device")) + t.Vprint("") + t.Vprintf(" Node: %s (%s)\n", reg.DisplayName, reg.ExternalNodeID) + t.Vprintf(" Linux user: %s\n", linuxUser.Username) + t.Vprint("") + + removed, err := sshcert.RemoveCertAuthorityLine(linuxUser.HomeDir, reg.ExternalNodeID, linuxUser.Username) + if err != nil { + return fmt.Errorf("disallow SSH failed: %w", err) + } + + if removed { + t.Vprint(t.Green(" Certificate authority removed from authorized_keys.")) + } else { + t.Vprint(t.Yellow(" No certificate authority line found in authorized_keys.")) + } + + // Legacy nodes store per-user keys instead of a cert-authority line. + // Remove them too; both operations are idempotent no-ops when nothing + // matches, so running both covers every node mode. + removedKeys, kerr := register.RemoveBrevAuthorizedKeys(linuxUser) + switch { + case kerr != nil: + t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: failed to remove Brev SSH keys: %v", kerr))) + case len(removedKeys) > 0: + t.Vprintf("%s Brev SSH keys removed from authorized_keys:\n", t.Green(" ✓")) + for _, key := range removedKeys { + t.Vprintf(" - %s\n", key) + } + } + + t.Vprint(t.Green("SSH disallowed. Run 'brev allow-ssh' to re-enable.")) + return nil +} diff --git a/pkg/cmd/enablessh/enablessh.go b/pkg/cmd/enablessh/enablessh.go deleted file mode 100644 index 9788b0e6f..000000000 --- a/pkg/cmd/enablessh/enablessh.go +++ /dev/null @@ -1,153 +0,0 @@ -// Package enablessh provides the brev enableSSH command for enabling SSH access -// to a registered external node. -package enablessh - -import ( - "context" - "fmt" - "os/exec" - "os/user" - - nodev1 "buf.build/gen/go/brevdev/devplane/protocolbuffers/go/devplaneapi/v1" - "connectrpc.com/connect" - - "github.com/brevdev/brev-cli/pkg/cmd/register" - "github.com/brevdev/brev-cli/pkg/config" - "github.com/brevdev/brev-cli/pkg/entity" - breverrors "github.com/brevdev/brev-cli/pkg/errors" - "github.com/brevdev/brev-cli/pkg/externalnode" - "github.com/brevdev/brev-cli/pkg/terminal" - - "github.com/spf13/cobra" -) - -// EnableSSHStore defines the store methods needed by the enableSSH command. -type EnableSSHStore interface { - GetCurrentUser() (*entity.User, error) - GetAccessToken() (string, error) -} - -// enableSSHDeps bundles the side-effecting dependencies of runEnableSSH so they -// can be replaced in tests. -type enableSSHDeps struct { - platform externalnode.PlatformChecker - nodeClients externalnode.NodeClientFactory - registrationStore register.RegistrationStore - prompter terminal.Selector -} - -func defaultEnableSSHDeps() enableSSHDeps { - return enableSSHDeps{ - platform: register.LinuxPlatform{}, - nodeClients: register.DefaultNodeClientFactory{}, - registrationStore: register.NewFileRegistrationStore(), - prompter: register.TerminalPrompter{}, - } -} - -func NewCmdEnableSSH(t *terminal.Terminal, store EnableSSHStore) *cobra.Command { - cmd := &cobra.Command{ - Annotations: map[string]string{"configuration": ""}, - Use: "enable-ssh", - DisableFlagsInUseLine: true, - Short: "Enable SSH access to this registered device", - Long: "Enable SSH access to this registered device for the current Brev user.", - Example: " brev enable-ssh", - RunE: func(cmd *cobra.Command, args []string) error { - return runEnableSSH(cmd.Context(), t, store, defaultEnableSSHDeps()) - }, - } - - return cmd -} - -func runEnableSSH(ctx context.Context, t *terminal.Terminal, s EnableSSHStore, deps enableSSHDeps) error { - if !deps.platform.IsCompatible() { - return fmt.Errorf("brev enable-ssh is only supported on Linux") - } - - reg, err := deps.registrationStore.Load() - if err != nil { - return fmt.Errorf("failed to read registration file: %w", err) - } - - brevUser, err := s.GetCurrentUser() - if err != nil { - return breverrors.WrapAndTrace(err) - } - - return enableSSH(ctx, t, deps, s, reg, brevUser) -} - -// enableSSH grants SSH access to the given node for the current Brev user. -// This is the "reflexive grant" — granting yourself SSH access to the device. -func enableSSH( - ctx context.Context, - t *terminal.Terminal, - deps enableSSHDeps, - tokenProvider externalnode.TokenProvider, - reg *register.DeviceRegistration, - brevUser *entity.User, -) error { - linuxUser, err := user.Current() - if err != nil { - return fmt.Errorf("failed to determine current Linux user: %w", err) - } - linuxUsername := linuxUser.Username - - checkSSHDaemon(t) - - t.Vprint("") - t.Vprint(t.Green("Enabling SSH access on this device")) - t.Vprint("") - t.Vprintf(" Node: %s (%s)\n", reg.DisplayName, reg.ExternalNodeID) - t.Vprintf(" Brev user: %s\n", brevUser.ID) - t.Vprintf(" Linux user: %s\n", linuxUsername) - t.Vprint("") - - node, err := fetchRegisteredNode(ctx, deps, tokenProvider, reg) - if err != nil { - return fmt.Errorf("enable SSH failed: %w", err) - } - - brevPortID, err := register.ResolveSSHAccessPort(ctx, t, deps.prompter, deps.nodeClients, tokenProvider, reg, node) - if err != nil { - return fmt.Errorf("enable SSH failed: %w", err) - } - - if err := register.SetupAndRegisterNodeSSHAccess(ctx, t, deps.nodeClients, tokenProvider, reg, brevUser, linuxUsername, brevPortID); err != nil { - return fmt.Errorf("enable SSH failed: %w", err) - } - - t.Vprint(t.Green(fmt.Sprintf("SSH access enabled. You can now SSH to this device via: brev shell %s", reg.DisplayName))) - return nil -} - -func fetchRegisteredNode( - ctx context.Context, - deps enableSSHDeps, - tokenProvider externalnode.TokenProvider, - reg *register.DeviceRegistration, -) (*nodev1.ExternalNode, error) { - client := deps.nodeClients.NewNodeClient(tokenProvider, config.GlobalConfig.GetBrevPublicAPIURL()) - resp, err := client.GetNode(ctx, connect.NewRequest(&nodev1.GetNodeRequest{ - ExternalNodeId: reg.ExternalNodeID, - OrganizationId: reg.OrgID, - })) - if err != nil { - return nil, fmt.Errorf("error retrieving node: %w", err) - } - return resp.Msg.GetExternalNode(), nil -} - -// checkSSHDaemon prints a warning if neither "ssh" nor "sshd" systemd services -// appear to be active. It never returns an error — it is best-effort. -func checkSSHDaemon(t *terminal.Terminal) { - for _, svc := range []string{"ssh", "sshd"} { - out, err := exec.Command("systemctl", "is-active", svc).Output() //nolint:gosec // fixed service names - if err == nil && len(out) > 0 && string(out[:len(out)-1]) == "active" { - return - } - } - t.Vprintf(" %s\n", t.Yellow("Warning: SSH daemon does not appear to be running. SSH access may not work until sshd is started.")) -} diff --git a/pkg/cmd/grantssh/grantssh.go b/pkg/cmd/grantssh/grantssh.go index 8d49a274a..7bf7981ab 100644 --- a/pkg/cmd/grantssh/grantssh.go +++ b/pkg/cmd/grantssh/grantssh.go @@ -184,7 +184,7 @@ func runGrantSSH(ctx context.Context, t *terminal.Terminal, s GrantSSHStore, opt } linuxUserOptions := uniqueLinuxUsersFromNodeSSHAccess(node) if len(linuxUserOptions) == 0 { - return fmt.Errorf("no Linux users on this node yet; run with --linux-user to specify one (e.g. after enable-ssh on the node)") + return fmt.Errorf("no Linux users on this node yet; run with --linux-user to specify one (e.g. after allow-ssh on the node)") } t.Vprint("") linuxUser = deps.prompter.Select("Select Linux user on the node", linuxUserOptions) diff --git a/pkg/cmd/register/register.go b/pkg/cmd/register/register.go index c83de152b..666e5d884 100644 --- a/pkg/cmd/register/register.go +++ b/pkg/cmd/register/register.go @@ -87,6 +87,8 @@ var ( registerLong = `Register your device with NVIDIA Brev This command registers this machine with Brev and brings up the Brev tunnel. +Registration does not enable SSH; run 'brev allow-ssh' afterwards to allow SSH +on this device, then 'brev grant-ssh' to grant users SSH access. Two modes are supported: • Interactive (default): run 'brev register' with no flags and follow prompts for device name and org. @@ -103,7 +105,10 @@ flow is used.` brev register # Non-interactive (--name and --org required) - brev register --name my-node --org my-org` + brev register --name my-node --org my-org + + # Allow SSH on this device after registering + brev allow-ssh` ) func NewCmdRegister(t *terminal.Terminal, store RegisterStore) *cobra.Command { @@ -319,6 +324,7 @@ func runRegisterSteps(ctx context.Context, t *terminal.Terminal, s RegisterStore Name: name, DeviceId: deviceID, NodeSpec: toProtoNodeSpec(hwProfile), + Labels: map[string]string{"sshprovider": "certauth"}, })) if err != nil { var connectErr *connect.Error diff --git a/pkg/sshcert/sshcert.go b/pkg/sshcert/sshcert.go index 84d8342bc..af24e5fa1 100644 --- a/pkg/sshcert/sshcert.go +++ b/pkg/sshcert/sshcert.go @@ -191,3 +191,58 @@ func writeAtomic(fs afero.Fs, path string, data []byte, mode os.FileMode) error } return breverrors.WrapAndTrace(fs.Rename(tmpName, path)) } + +// CertAuthorityPrincipal returns the SSH certificate principal for a node and +// Linux user +func CertAuthorityPrincipal(nodeID, linuxUser string) string { + return fmt.Sprintf("brev:v1:vm:%s:login:%s", nodeID, linuxUser) +} + +// RemoveCertAuthorityLine removes the cert-authority line for the given node +// and Linux user from ~/.ssh/authorized_keys. Returns true if a line was +// removed. Missing file is treated as nothing-to-remove. +func RemoveCertAuthorityLine(homeDir, nodeID, linuxUser string) (bool, error) { + prefix := fmt.Sprintf("cert-authority,principals=%q ", CertAuthorityPrincipal(nodeID, linuxUser)) + + authKeysPath := filepath.Join(homeDir, ".ssh", "authorized_keys") + + existing, err := os.ReadFile(authKeysPath) // #nosec G304 + if err != nil { + if os.IsNotExist(err) { + return false, nil + } + return false, fmt.Errorf("reading authorized_keys: %w", err) + } + + var kept []string + var removed bool + for line := range strings.SplitSeq(string(existing), "\n") { + trimmed := strings.TrimSpace(line) + if strings.HasPrefix(trimmed, prefix) && strings.Contains(trimmed, "cert-authority") { + removed = true + continue + } + kept = append(kept, line) + } + + if !removed { + return false, nil + } + + result := strings.Join(kept, "\n") + if err := os.WriteFile(authKeysPath, []byte(result), 0o600); err != nil { + return false, fmt.Errorf("writing authorized_keys: %w", err) + } + + return true, nil +} + +// HasCertAuthorityLine reports whether authorized_keys contains a Brev +// cert-authority line for the given node (any Linux user). +func HasCertAuthorityLine(homeDir, nodeID string) bool { + data, err := os.ReadFile(filepath.Join(homeDir, ".ssh", "authorized_keys")) // #nosec G304 + if err != nil { + return false + } + return strings.Contains(string(data), "brev:v1:vm:"+nodeID+":") +} diff --git a/pkg/sshcert/sshcert_test.go b/pkg/sshcert/sshcert_test.go index 5e7299b04..f25ea83ea 100644 --- a/pkg/sshcert/sshcert_test.go +++ b/pkg/sshcert/sshcert_test.go @@ -4,6 +4,9 @@ import ( "crypto/ed25519" "crypto/rand" "encoding/pem" + "os" + "os/user" + "path/filepath" "strings" "testing" "time" @@ -234,3 +237,131 @@ func mustGen(t *testing.T) ([]byte, string) { } return priv, pub } + +func testHomeDir(t *testing.T) *user.User { + t.Helper() + return &user.User{HomeDir: t.TempDir()} +} + +func Test_RemoveCertAuthorityLine_RemovesMatchingLine(t *testing.T) { + u := testHomeDir(t) + sshDir := filepath.Join(u.HomeDir, ".ssh") + if err := os.MkdirAll(sshDir, 0o700); err != nil { + t.Fatal(err) + } + + caKey := "ssh-ed25519 AAAAC3Nz dummyCA" + entry := `cert-authority,principals="brev:v1:vm:unode_abc:login:ubuntu" ` + caKey + content := strings.Join([]string{ + "ssh-rsa EXISTING user@host", + entry, + "ssh-ed25519 OTHER admin@server", + "", + }, "\n") + + if err := os.WriteFile(filepath.Join(sshDir, "authorized_keys"), []byte(content), 0o600); err != nil { + t.Fatal(err) + } + + removed, err := RemoveCertAuthorityLine(u.HomeDir, "unode_abc", "ubuntu") + if err != nil { + t.Fatalf("removeCertAuthority: %v", err) + } + if !removed { + t.Fatal("expected line to be removed") + } + + data, err := os.ReadFile(filepath.Join(sshDir, "authorized_keys")) + if err != nil { + t.Fatal(err) + } + + result := string(data) + if strings.Contains(result, caKey) { + t.Errorf("CA key still present:\n%s", result) + } + if strings.Contains(result, "cert-authority") { + t.Errorf("cert-authority line still present:\n%s", result) + } + if !strings.Contains(result, "ssh-rsa EXISTING user@host") { + t.Errorf("non-brev key was removed:\n%s", result) + } +} + +func Test_RemoveCertAuthorityLine_NoopWhenFileDoesNotExist(t *testing.T) { + u := testHomeDir(t) + removed, err := RemoveCertAuthorityLine(u.HomeDir, "unode_abc", "ubuntu") + if err != nil { + t.Fatalf("expected no error for missing file: %v", err) + } + if removed { + t.Error("expected removed=false for missing file") + } +} + +func Test_RemoveCertAuthorityLine_NoopWhenNoMatch(t *testing.T) { + u := testHomeDir(t) + sshDir := filepath.Join(u.HomeDir, ".ssh") + if err := os.MkdirAll(sshDir, 0o700); err != nil { + t.Fatal(err) + } + + original := "ssh-rsa EXISTING user@host\n" + if err := os.WriteFile(filepath.Join(sshDir, "authorized_keys"), []byte(original), 0o600); err != nil { + t.Fatal(err) + } + + removed, err := RemoveCertAuthorityLine(u.HomeDir, "unode_abc", "ubuntu") + if err != nil { + t.Fatalf("removeCertAuthority: %v", err) + } + if removed { + t.Error("expected removed=false when no match") + } + + data, err := os.ReadFile(filepath.Join(sshDir, "authorized_keys")) + if err != nil { + t.Fatal(err) + } + if string(data) != original { + t.Errorf("file was modified when it shouldn't have been") + } +} + +func Test_RemoveCertAuthorityLine_OnlyRemovesMatchingPrincipal(t *testing.T) { + u := testHomeDir(t) + sshDir := filepath.Join(u.HomeDir, ".ssh") + if err := os.MkdirAll(sshDir, 0o700); err != nil { + t.Fatal(err) + } + + otherEntry := `cert-authority,principals="brev:v1:vm:other_node:login:ubuntu" ssh-ed25519 OTHER_CA` + targetEntry := `cert-authority,principals="brev:v1:vm:unode_abc:login:ubuntu" ssh-ed25519 TARGET_CA` + content := strings.Join([]string{ + otherEntry, + targetEntry, + "", + }, "\n") + + if err := os.WriteFile(filepath.Join(sshDir, "authorized_keys"), []byte(content), 0o600); err != nil { + t.Fatal(err) + } + + removed, err := RemoveCertAuthorityLine(u.HomeDir, "unode_abc", "ubuntu") + if err != nil { + t.Fatalf("removeCertAuthority: %v", err) + } + if !removed { + t.Fatal("expected line to be removed") + } + + data, _ := os.ReadFile(filepath.Join(sshDir, "authorized_keys")) + result := string(data) + + if strings.Contains(result, "TARGET_CA") { + t.Errorf("target CA still present:\n%s", result) + } + if !strings.Contains(result, "OTHER_CA") { + t.Errorf("other node's CA was removed:\n%s", result) + } +}