From 1bb166b8581664b14e4537c8daab28f271003973 Mon Sep 17 00:00:00 2001 From: Pratik Patel Date: Wed, 19 Aug 2026 16:54:02 -0700 Subject: [PATCH 1/2] simplify register --- pkg/auth/auth.go | 29 +- pkg/auth/auth_test.go | 90 +- pkg/cmd/cmd.go | 54 +- pkg/cmd/cmd_test.go | 67 -- pkg/cmd/deregister/deregister.go | 62 +- pkg/cmd/deregister/deregister_test.go | 186 +++- pkg/cmd/enablessh/enablessh.go | 2 +- pkg/cmd/grantssh/grantssh_test.go | 2 +- pkg/cmd/login/login.go | 37 +- pkg/cmd/login/login_test.go | 117 ++- pkg/cmd/ls/ls_test.go | 2 +- pkg/cmd/register/device_registration_store.go | 38 +- .../device_registration_store_test.go | 69 +- pkg/cmd/register/register.go | 313 ++++--- pkg/cmd/register/register_test.go | 874 +++++++++--------- pkg/cmd/register/sshkeys.go | 2 +- pkg/cmd/revokessh/revokessh_test.go | 2 +- 17 files changed, 1189 insertions(+), 757 deletions(-) diff --git a/pkg/auth/auth.go b/pkg/auth/auth.go index 21de0d85..d79f30c2 100644 --- a/pkg/auth/auth.go +++ b/pkg/auth/auth.go @@ -102,7 +102,9 @@ type Auth struct { const BrevAPIKeyPrefix = "bak-" -const MissingAPIKeyOrgIDMessage = "api key auth requires an org id; run brev login --api-key --org-id " +const AccessKeyEnvVar = "BREV_ACCESS_KEY" + +const MissingAPIKeyOrgIDMessage = "org id missing, please login again; run 'brev login --api-key '" type APIKeyAuthStore interface { GetAuthTokens() (*entity.AuthTokens, error) @@ -142,6 +144,9 @@ func IsBrevAPIKey(token string) bool { } func IsAPIKeyAuthStore(authTokensProvider APIKeyAuthStore) bool { + if strings.TrimSpace(os.Getenv(AccessKeyEnvVar)) != "" { + return true + } tokens, err := authTokensProvider.GetAuthTokens() if err != nil { return false @@ -153,6 +158,20 @@ func IsAPIKeyAuthStore(authTokensProvider APIKeyAuthStore) bool { } func GetAPIKeyOrgID(authTokensProvider APIKeyAuthStore) (string, error) { + if envKey := strings.TrimSpace(os.Getenv(AccessKeyEnvVar)); envKey != "" { + tokens, err := authTokensProvider.GetAuthTokens() + if err != nil { + return "", breverrors.WrapAndTrace(err) + } + if tokens == nil || tokens.APIKey != envKey { + return "", breverrors.NewValidationError(MissingAPIKeyOrgIDMessage) + } + orgID := strings.TrimSpace(tokens.APIKeyOrgID) + if orgID == "" { + return "", breverrors.NewValidationError(MissingAPIKeyOrgIDMessage) + } + return orgID, nil + } tokens, err := authTokensProvider.GetAuthTokens() if err != nil { return "", breverrors.WrapAndTrace(err) @@ -204,6 +223,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(AccessKeyEnvVar)); key != "" { + return key, nil + } tokens, err := t.getSavedTokensOrNil() if err != nil { return "", breverrors.WrapAndTrace(err) @@ -217,7 +239,6 @@ func (t Auth) GetFreshAccessTokenOrNil() (string, error) { return apiKey, nil } - // should always at least have access token? if tokens.AccessToken == "" { breverrors.GetDefaultErrorReporter().ReportMessage("access token is an empty string but shouldn't be") } @@ -301,10 +322,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 4c271b0b..88444478 100644 --- a/pkg/auth/auth_test.go +++ b/pkg/auth/auth_test.go @@ -28,15 +28,6 @@ func TestIsAccessTokenValid(t *testing.T) { if !assert.False(t, res) { return } - - // expiredToken := "eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCIsImtpZCI6ImdLTXBESXlRc0ZXSF9zYWdiT2oyViJ9.eyJpc3MiOiJodHRwczovL2JyZXZkZXYudXMuYXV0aDAuY29tLyIsInN1YiI6Imdvb2dsZS1vYXV0aDJ8MTAxNzY0NjMwNTEwODYxNDk5MTgwIiwiYXVkIjpbImh0dHBzOi8vYnJldmRldi51cy5hdXRoMC5jb20vYXBpL3YyLyIsImh0dHBzOi8vYnJldmRldi51cy5hdXRoMC5jb20vdXNlcmluZm8iXSwiaWF0IjoxNjM4NTYyMzY4LCJleHAiOjE2Mzg2NDg3NjgsImF6cCI6IkphcUpSTEVzZGF0NXc3VGIwV3FtVHh6SWVxd3FlcG1rIiwic2NvcGUiOiJvcGVuaWQgcHJvZmlsZSBlbWFpbCBvZmZsaW5lX2FjY2VzcyJ9.YCiO-som26ehT91qGAX5ZfrtVg4eYwamnlMRoCuUljXmg8Nf-ArDyoG32CqZkQ6YJ5XnzrVX9bVk5ZNHP_AFSE9SJvYL6MchoN09nR84WTbevRBCtZedIZUk5ULg6rWo5mszGr-S2gi08od4iTzXtKySPx1JnT60muRj_k9VV3MyixqvngEz5NvmFDdA8glGes5_iOuiBidmjOJzi_CVfKJ9s48BhlxzciSXFC0_DUBnT9OThjYjUP-22ohOuWwJWomRUv6gMSq78hJOALc330LwvmEsLdzlP7a3otIYM43hTtAVJ9QEL6M08GKqm3PdikzTxiGdfuQUhgMDlXygbQ" - // res, err := isAccessTokenValid(expiredToken) - // if !assert.Nil(t, err) { - // return - // } - // if !assert.False(t, res) { - // return - // } } type MockAuthStore struct { @@ -48,6 +39,7 @@ type MockAuthStore struct { func (m *MockAuthStore) SaveAuthTokens(tokens entity.AuthTokens) error { m.saved = tokens m.didSave = true + m.authTokens = &tokens // write-then-read consistent (mirrors a real store) return nil } @@ -110,6 +102,7 @@ func (s *sideEffectingTokenStore) GetAccessToken() (string, error) { } func TestIsAPIKeyAuthStore_ReadsSavedTokensWithoutAccessTokenSideEffects(t *testing.T) { + t.Setenv(AccessKeyEnvVar, "") s := &sideEffectingTokenStore{ tokens: &entity.AuthTokens{APIKey: testAPIKey}, } @@ -119,6 +112,7 @@ func TestIsAPIKeyAuthStore_ReadsSavedTokensWithoutAccessTokenSideEffects(t *test } func TestIsAPIKeyAuthStore_LegacyCredentialsAreNotAPIKeyAuth(t *testing.T) { + t.Setenv(AccessKeyEnvVar, "") s := &sideEffectingTokenStore{ tokens: &entity.AuthTokens{ AccessToken: validToken, @@ -130,6 +124,35 @@ func TestIsAPIKeyAuthStore_LegacyCredentialsAreNotAPIKeyAuth(t *testing.T) { assert.False(t, s.getAccessTokenCalled) } +func TestIsAPIKeyAuthStore_EnvKeyIsAPIKeyEvenWhenNotPersisted(t *testing.T) { + t.Setenv(AccessKeyEnvVar, testAPIKey) + s := &sideEffectingTokenStore{tokens: nil} // nothing persisted + assert.True(t, IsAPIKeyAuthStore(s)) +} + +func TestGetAPIKeyOrgID_EnvKeyMismatchingPersistedRejects(t *testing.T) { + t.Setenv(AccessKeyEnvVar, BrevAPIKeyPrefix+"env-key") + s := &sideEffectingTokenStore{tokens: &entity.AuthTokens{ + APIKey: BrevAPIKeyPrefix + "persisted-key", + APIKeyOrgID: "org-persisted", + }} + _, err := GetAPIKeyOrgID(s) + assert.Error(t, err) + assert.Contains(t, err.Error(), "org id missing") +} + +// When the env key matches the persisted key, its persisted org is valid. +func TestGetAPIKeyOrgID_EnvKeyMatchingPersistedReturnsOrg(t *testing.T) { + t.Setenv(AccessKeyEnvVar, testAPIKey) + s := &sideEffectingTokenStore{tokens: &entity.AuthTokens{ + APIKey: testAPIKey, + APIKeyOrgID: "org-test", + }} + orgID, err := GetAPIKeyOrgID(s) + assert.NoError(t, err) + assert.Equal(t, "org-test", orgID) +} + type cliAuthStore struct { tokens *entity.AuthTokens user *entity.User @@ -238,6 +261,43 @@ func TestGetFreshAccessTokenOrNil_APIKeyOnlyCredentialReturnsAPIKey(t *testing.T assert.False(t, s.didSave) } +// Closest credential wins: BREV_ACCESS_KEY takes precedence over saved +// tokens (flag/env before persisted), matching other CLIs. The global +// --api-key flag handler populates this env var before the auth chain runs. +func TestGetFreshAccessTokenOrNil_EnvVarTakesPrecedenceOverSaved(t *testing.T) { + t.Setenv(AccessKeyEnvVar, 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_ACCESS_KEY must win over saved tokens") +} + +// With no saved credential, BREV_ACCESS_KEY authenticates headless/CI commands. +func TestGetFreshAccessTokenOrNil_EnvVarFallbackWhenNoSavedTokens(t *testing.T) { + t.Setenv(AccessKeyEnvVar, 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(AccessKeyEnvVar, "") + 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{ @@ -286,18 +346,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() { diff --git a/pkg/cmd/cmd.go b/pkg/cmd/cmd.go index 8aa4c561..b87045c4 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" @@ -59,7 +60,6 @@ import ( "github.com/brevdev/brev-cli/pkg/cmd/upgrade" "github.com/brevdev/brev-cli/pkg/cmd/version" "github.com/brevdev/brev-cli/pkg/config" - "github.com/brevdev/brev-cli/pkg/entity" "github.com/brevdev/brev-cli/pkg/featureflag" "github.com/brevdev/brev-cli/pkg/files" "github.com/brevdev/brev-cli/pkg/remoteversion" @@ -73,6 +73,7 @@ import ( var ( userFlag string + apiKeyFlag string printVersion bool noCheckLatest bool ) @@ -84,6 +85,8 @@ 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") 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 +166,9 @@ func NewBrevCommand() *cobra.Command { //nolint:funlen,gocognit,gocyclo // defin fmt.Println(v) } } + if apiKeyFlag != "" { + os.Setenv(auth.AccessKeyEnvVar, apiKeyFlag) + } if userFlag != "" { _, err := noLoginCmdStore.WithUserID(userFlag) if err != nil { @@ -233,33 +239,21 @@ func NewBrevCommand() *cobra.Command { //nolint:funlen,gocognit,gocyclo // defin cmds.SetUsageTemplate(usageTemplate) - // In-memory auth for external node commands — never touches credentials.json. - // Pre-fill the cached email so the user sees a confirmation prompt instead of - // having to type it from scratch every time. - cachedEmail, _ := fsStore.GetCachedEmail() - memAuthenticator := auth.StandardLogin("", cachedEmail, nil) - if cachedEmail != "" { - if kas, ok := memAuthenticator.(auth.KasAuthenticator); ok { - kas.ShouldPromptEmail = true - memAuthenticator = kas - } - } - memAuthStore := &emailCachingAuthStore{ - MemoryAuthStore: store.NewMemoryAuthStore(), - fileStore: fsStore, - } - memLoginAuth := auth.NewLoginAuth(memAuthStore, memAuthenticator) - memLoginAuth.WithShouldLogin(func() (bool, error) { return true, nil }) - + // External node commands (register/deregister/enable-ssh/grant-ssh/revoke-ssh) + // read credentials.json and BREV_ACCESS_KEY but never prompt for a login — + // a shared box should not be encouraged to write durable creds. externalNodeCmdStore := fsStore.WithNoAuthHTTPClient( store.NewNoAuthHTTPClient(conf.GetBrevAPIURl()), - ).WithAuth(memLoginAuth, store.WithDebug(conf.GetDebugHTTP())) + ).WithAuth(noLoginAuth, store.WithDebug(conf.GetDebugHTTP())) err = externalNodeCmdStore.SetForbiddenStatusRetryHandler(func() error { - _, err1 := memLoginAuth.GetAccessToken() + token, err1 := noLoginAuth.GetAccessToken() if err1 != nil { return breverrors.WrapAndTrace(err1) } + if token == "" { + return breverrors.New("not authenticated; set BREV_ACCESS_KEY or run 'brev login --api-key' on a trusted machine") + } return nil }) if err != nil { @@ -545,22 +539,4 @@ var ( _ store.Auth = auth.NoLoginAuth{} _ auth.AuthStore = store.FileStore{} _ auth.AuthStore = &store.MemoryAuthStore{} - _ auth.AuthStore = &emailCachingAuthStore{} ) - -// emailCachingAuthStore wraps MemoryAuthStore and persists the login email -// to ~/.brev/cached-email after each successful authentication. -type emailCachingAuthStore struct { - *store.MemoryAuthStore - fileStore *store.FileStore -} - -func (e *emailCachingAuthStore) SaveAuthTokens(tokens entity.AuthTokens) error { - if err := e.MemoryAuthStore.SaveAuthTokens(tokens); err != nil { - return breverrors.WrapAndTrace(err) - } - if email := auth.GetEmailFromToken(tokens.AccessToken); email != "" { - _ = e.fileStore.SaveCachedEmail(email) - } - return nil -} diff --git a/pkg/cmd/cmd_test.go b/pkg/cmd/cmd_test.go index 5f748637..6c2487e7 100644 --- a/pkg/cmd/cmd_test.go +++ b/pkg/cmd/cmd_test.go @@ -1,37 +1,13 @@ package cmd import ( - "encoding/base64" - "encoding/json" "testing" - "github.com/brevdev/brev-cli/pkg/entity" - "github.com/brevdev/brev-cli/pkg/store" - "github.com/spf13/afero" "github.com/spf13/cobra" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) -// fakeJWT builds an unsigned JWT with the given claims (header.payload.signature). -func fakeJWT(t *testing.T, claims map[string]interface{}) string { - t.Helper() - header := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"none","typ":"JWT"}`)) - payload, err := json.Marshal(claims) - require.NoError(t, err) - return header + "." + base64.RawURLEncoding.EncodeToString(payload) + "." -} - -func newTestFileStore(t *testing.T) *store.FileStore { - t.Helper() - fs := afero.NewMemMapFs() - err := fs.MkdirAll("/home/testuser/.brev", 0o755) - require.NoError(t, err) - return store.NewBasicStore().WithFileSystem(fs).WithUserHomeDirGetter( - func() (string, error) { return "/home/testuser", nil }, - ) -} - func TestAccessCommandsExcludesHiddenCommands(t *testing.T) { root := &cobra.Command{Use: "brev"} visible := &cobra.Command{ @@ -51,46 +27,3 @@ func TestAccessCommandsExcludesHiddenCommands(t *testing.T) { require.Len(t, commands, 1) assert.Same(t, visible, commands[0]) } - -func TestEmailCachingAuthStore_SaveCachesEmail(t *testing.T) { - fs := newTestFileStore(t) - s := &emailCachingAuthStore{ - MemoryAuthStore: store.NewMemoryAuthStore(), - fileStore: fs, - } - - token := fakeJWT(t, map[string]interface{}{"email": "user@example.com"}) - err := s.SaveAuthTokens(entity.AuthTokens{AccessToken: token}) - require.NoError(t, err) - - cached, err := fs.GetCachedEmail() - require.NoError(t, err) - assert.Equal(t, "user@example.com", cached) -} - -func TestEmailCachingAuthStore_NoEmailInToken(t *testing.T) { - fs := newTestFileStore(t) - s := &emailCachingAuthStore{ - MemoryAuthStore: store.NewMemoryAuthStore(), - fileStore: fs, - } - - token := fakeJWT(t, map[string]interface{}{"sub": "12345"}) - err := s.SaveAuthTokens(entity.AuthTokens{AccessToken: token}) - require.NoError(t, err) - - cached, err := fs.GetCachedEmail() - require.NoError(t, err) - assert.Equal(t, "", cached) -} - -func TestEmailCachingAuthStore_EmptyAccessToken(t *testing.T) { - fs := newTestFileStore(t) - s := &emailCachingAuthStore{ - MemoryAuthStore: store.NewMemoryAuthStore(), - fileStore: fs, - } - - err := s.SaveAuthTokens(entity.AuthTokens{AccessToken: ""}) - require.Error(t, err) -} diff --git a/pkg/cmd/deregister/deregister.go b/pkg/cmd/deregister/deregister.go index efd9090b..5bcdac5c 100644 --- a/pkg/cmd/deregister/deregister.go +++ b/pkg/cmd/deregister/deregister.go @@ -3,6 +3,7 @@ package deregister import ( "context" + "errors" "fmt" "os/user" @@ -20,18 +21,15 @@ import ( "github.com/spf13/cobra" ) -// DeregisterStore defines the store methods needed by the deregister command. type DeregisterStore interface { GetCurrentUser() (*entity.User, error) GetAccessToken() (string, error) } -// SSHKeyRemover removes Brev-managed SSH keys and returns the lines removed. type SSHKeyRemover interface { RemoveBrevKeys(u *user.User) ([]string, error) } -// brevSSHKeyRemover delegates to register.RemoveBrevAuthorizedKeys. type brevSSHKeyRemover struct{} func (brevSSHKeyRemover) RemoveBrevKeys(u *user.User) ([]string, error) { @@ -97,6 +95,53 @@ func NewCmdDeregister(t *terminal.Terminal, store DeregisterStore) *cobra.Comman return cmd } +func removeNodeFromBrev(ctx context.Context, t *terminal.Terminal, s DeregisterStore, deps deregisterDeps, reg *register.DeviceRegistration) error { + externalNodeID := reg.ExternalNodeID + if externalNodeID == "" && reg.DeviceID != "" { + lookedUp, lookupErr := findNodeByDeviceID(ctx, s, deps, reg.OrgID, reg.DeviceID) + if lookupErr != nil { + t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Could not look up pending node by device ID: %v", lookupErr))) + } + if lookedUp != "" { + externalNodeID = lookedUp + } + } + if externalNodeID == "" { + t.Vprintf(" %s\n", t.Yellow("No registered node to remove (pending registration); cleaning up local state.")) + return nil + } + client := deps.nodeClients.NewNodeClient(s, config.GlobalConfig.GetBrevPublicAPIURL()) + _, err := client.RemoveNode(ctx, connect.NewRequest(&nodev1.RemoveNodeRequest{ + ExternalNodeId: externalNodeID, + })) + if err != nil { + var connectErr *connect.Error + if errors.As(err, &connectErr) && connectErr.Code() == connect.CodeNotFound { + t.Vprintf(" %s\n", t.Yellow("Node not found on Brev; continuing.")) + return nil + } + return fmt.Errorf("failed to deregister node: %w", err) + } + t.Vprintf("%s Node removed from Brev.\n", t.Green(" ✓")) + return nil +} + +func findNodeByDeviceID(ctx context.Context, s externalnode.TokenProvider, deps deregisterDeps, orgID, deviceID string) (string, error) { + client := deps.nodeClients.NewNodeClient(s, config.GlobalConfig.GetBrevPublicAPIURL()) + resp, err := client.ListNodes(ctx, connect.NewRequest(&nodev1.ListNodesRequest{ + OrganizationId: orgID, + })) + if err != nil { + return "", fmt.Errorf("failed to list nodes: %w", err) + } + for _, n := range resp.Msg.GetItems() { + if n.GetDeviceId() == deviceID { + return n.GetExternalNodeId(), nil + } + } + return "", nil +} + func runDeregister(ctx context.Context, t *terminal.Terminal, s DeregisterStore, deps deregisterDeps, skipConfirm bool) error { //nolint:funlen,gocyclo // deregistration flow if !deps.platform.IsCompatible() { return fmt.Errorf("brev deregister is only supported on Linux") @@ -106,7 +151,7 @@ func runDeregister(ctx context.Context, t *terminal.Terminal, s DeregisterStore, return fmt.Errorf("sudo issue: %w", err) } - reg, err := deps.registrationStore.Load() + reg, err := deps.registrationStore.Load(true) // deregister should still work for pending registrations if err != nil { return err //nolint:wrapcheck // do not present stack trace for this error } @@ -158,14 +203,9 @@ func runDeregister(ctx context.Context, t *terminal.Terminal, s DeregisterStore, } t.Vprint(t.Yellow("[Step 1/4] Removing node from Brev...")) - client := deps.nodeClients.NewNodeClient(s, config.GlobalConfig.GetBrevPublicAPIURL()) - _, err = client.RemoveNode(ctx, connect.NewRequest(&nodev1.RemoveNodeRequest{ - ExternalNodeId: reg.ExternalNodeID, - })) - if err != nil { - return fmt.Errorf("failed to deregister node: %w", err) + if err := removeNodeFromBrev(ctx, t, s, deps, reg); err != nil { + return err } - t.Vprintf("%s Node removed from Brev.\n", t.Green(" ✓")) t.Vprint("") t.Vprint(t.Yellow("[Step 2/4] Removing Brev SSH keys...")) diff --git a/pkg/cmd/deregister/deregister_test.go b/pkg/cmd/deregister/deregister_test.go index 95c5ac99..2389ccd3 100644 --- a/pkg/cmd/deregister/deregister_test.go +++ b/pkg/cmd/deregister/deregister_test.go @@ -37,6 +37,7 @@ func (m *mockDeregisterStore) GetAccessToken() (string, error) { return m.token, type fakeNodeService struct { nodev1connect.UnimplementedExternalNodeServiceHandler removeNodeFn func(*nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) + listNodesFn func(*nodev1.ListNodesRequest) (*nodev1.ListNodesResponse, error) } func (f *fakeNodeService) RemoveNode(_ context.Context, req *connect.Request[nodev1.RemoveNodeRequest]) (*connect.Response[nodev1.RemoveNodeResponse], error) { @@ -47,6 +48,14 @@ func (f *fakeNodeService) RemoveNode(_ context.Context, req *connect.Request[nod return connect.NewResponse(resp), nil } +func (f *fakeNodeService) ListNodes(_ context.Context, req *connect.Request[nodev1.ListNodesRequest]) (*connect.Response[nodev1.ListNodesResponse], error) { + resp, err := f.listNodesFn(req.Msg) + if err != nil { + return nil, err + } + return connect.NewResponse(resp), nil +} + // mockRegistrationStore satisfies register.RegistrationStore for deregister tests. type mockRegistrationStore struct { reg *register.DeviceRegistration @@ -57,7 +66,7 @@ func (m *mockRegistrationStore) Save(reg *register.DeviceRegistration) error { return nil } -func (m *mockRegistrationStore) Load() (*register.DeviceRegistration, error) { +func (m *mockRegistrationStore) Load(bool) (*register.DeviceRegistration, error) { if m.reg == nil { return nil, fmt.Errorf("no registration") } @@ -281,7 +290,6 @@ func Test_runDeregister_RemoveNodeFails(t *testing.T) { t.Fatal("expected error when RemoveNode fails") } - // Registration should still exist (server-side removal failed) exists, err := regStore.Exists() if err != nil { t.Fatalf("Exists error: %v", err) @@ -291,6 +299,180 @@ 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", + } + + svc := &fakeNodeService{ + removeNodeFn: func(_ *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { + return nil, connect.NewError(connect.CodeNotFound, nil) + }, + } + + deps, server := testDeregisterDeps(t, svc, regStore) + defer server.Close() + + term := terminal.New() + err := runDeregister(context.Background(), term, store, deps, false) + if err != nil { + t.Fatalf("NotFound should be treated as success (node already gone), got: %v", err) + } + + exists, err := regStore.Exists() + if err != nil { + t.Fatalf("Exists error: %v", err) + } + if exists { + t.Error("expected local registration to be deleted even when RemoveNode returns NotFound") + } +} + +func Test_runDeregister_PendingRegistration_SkipsRemoveNodeAndCleansUp(t *testing.T) { + regStore := &mockRegistrationStore{ + reg: ®ister.DeviceRegistration{ + DisplayName: "My Spark", + OrgID: "org_123", + DeviceID: "dev-uuid-pending", + Status: register.RegistrationStatusPending, + }, + } + + store := &mockDeregisterStore{ + user: &entity.User{ID: "user_1"}, + token: "tok", + } + + var removeCalled bool + svc := &fakeNodeService{ + removeNodeFn: func(req *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { + removeCalled = true + return nil, fmt.Errorf("RemoveNode should not be called with empty ID %q", req.GetExternalNodeId()) + }, + listNodesFn: func(*nodev1.ListNodesRequest) (*nodev1.ListNodesResponse, error) { + return &nodev1.ListNodesResponse{}, nil // no backend node matches the device ID + }, + } + + deps, server := testDeregisterDeps(t, svc, regStore) + defer server.Close() + + term := terminal.New() + err := runDeregister(context.Background(), term, store, deps, false) + if err != nil { + t.Fatalf("deregister of a pending registration should succeed, got: %v", err) + } + + if removeCalled { + t.Error("RemoveNode should not be called for a pending registration (no ExternalNodeID)") + } + + exists, err := regStore.Exists() + if err != nil { + t.Fatalf("Exists error: %v", err) + } + if exists { + t.Error("expected local pending registration to be deleted") + } +} + +// Test_runDeregister_PendingRegistration_LooksUpAndRemovesBackendNode verifies that +// when a pending record has no ExternalNodeID but AddNode succeeded backend-side, +// deregister recovers the node via ListNodes (by device ID) and removes it. +func Test_runDeregister_PendingRegistration_LooksUpAndRemovesBackendNode(t *testing.T) { + const deviceID = "dev-uuid-pending" + const externalNodeID = "unode_recovered" + regStore := &mockRegistrationStore{ + reg: ®ister.DeviceRegistration{ + DisplayName: "My Spark", + OrgID: "org_123", + DeviceID: deviceID, + Status: register.RegistrationStatusPending, + }, + } + + store := &mockDeregisterStore{user: &entity.User{ID: "user_1"}, token: "tok"} + + var gotOrgID string + var removedNodeID string + svc := &fakeNodeService{ + listNodesFn: func(req *nodev1.ListNodesRequest) (*nodev1.ListNodesResponse, error) { + gotOrgID = req.GetOrganizationId() + return &nodev1.ListNodesResponse{ + Items: []*nodev1.ExternalNode{ + {ExternalNodeId: "unode_other", DeviceId: "dev-different"}, + {ExternalNodeId: externalNodeID, DeviceId: deviceID}, + }, + }, nil + }, + removeNodeFn: func(req *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { + removedNodeID = req.GetExternalNodeId() + return &nodev1.RemoveNodeResponse{}, nil + }, + } + + deps, server := testDeregisterDeps(t, svc, regStore) + defer server.Close() + + term := terminal.New() + if err := runDeregister(context.Background(), term, store, deps, false); err != nil { + t.Fatalf("deregister failed: %v", err) + } + + if gotOrgID != "org_123" { + t.Errorf("expected ListNodes scoped to org_123, got %q", gotOrgID) + } + if removedNodeID != externalNodeID { + t.Errorf("expected RemoveNode called with recovered ID %q, got %q", externalNodeID, removedNodeID) + } + exists, _ := regStore.Exists() + if exists { + t.Error("expected local registration to be deleted") + } +} + +// Test_runDeregister_PendingRegistration_ListNodesFails verifies that a +// ListNodes failure during the device-ID lookup is non-fatal: deregister still +// cleans up local state (the backend node, if any, can be deleted in the UI). +func Test_runDeregister_PendingRegistration_ListNodesFails(t *testing.T) { + regStore := &mockRegistrationStore{ + reg: ®ister.DeviceRegistration{ + DisplayName: "My Spark", + OrgID: "org_123", + DeviceID: "dev-uuid-pending", + Status: register.RegistrationStatusPending, + }, + } + store := &mockDeregisterStore{user: &entity.User{ID: "user_1"}, token: "tok"} + + svc := &fakeNodeService{ + listNodesFn: func(*nodev1.ListNodesRequest) (*nodev1.ListNodesResponse, error) { + return nil, connect.NewError(connect.CodeInternal, nil) + }, + } + + deps, server := testDeregisterDeps(t, svc, regStore) + defer server.Close() + + term := terminal.New() + if err := runDeregister(context.Background(), term, store, deps, false); err != nil { + t.Fatalf("ListNodes failure should be non-fatal, got: %v", err) + } + exists, _ := regStore.Exists() + if exists { + t.Error("expected local registration to be deleted despite ListNodes failure") + } +} + func Test_runDeregister_AlwaysUninstallsNetbird(t *testing.T) { regStore := &mockRegistrationStore{ reg: ®ister.DeviceRegistration{ diff --git a/pkg/cmd/enablessh/enablessh.go b/pkg/cmd/enablessh/enablessh.go index 9788b0e6..6b6ac548 100644 --- a/pkg/cmd/enablessh/enablessh.go +++ b/pkg/cmd/enablessh/enablessh.go @@ -66,7 +66,7 @@ func runEnableSSH(ctx context.Context, t *terminal.Terminal, s EnableSSHStore, d return fmt.Errorf("brev enable-ssh is only supported on Linux") } - reg, err := deps.registrationStore.Load() + reg, err := deps.registrationStore.Load(false) if err != nil { return fmt.Errorf("failed to read registration file: %w", err) } diff --git a/pkg/cmd/grantssh/grantssh_test.go b/pkg/cmd/grantssh/grantssh_test.go index 10b3cbab..e1e6441a 100644 --- a/pkg/cmd/grantssh/grantssh_test.go +++ b/pkg/cmd/grantssh/grantssh_test.go @@ -44,7 +44,7 @@ func (m *mockRegistrationStore) Save(reg *register.DeviceRegistration) error { return nil } -func (m *mockRegistrationStore) Load() (*register.DeviceRegistration, error) { +func (m *mockRegistrationStore) Load(bool) (*register.DeviceRegistration, error) { if m.reg == nil { return nil, fmt.Errorf("no registration") } diff --git a/pkg/cmd/login/login.go b/pkg/cmd/login/login.go index 51d044c4..2eb0cc0c 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 @@ -103,9 +105,9 @@ func NewCmdLogin(t *terminal.Terminal, loginStore LoginStore, auth Auth) *cobra. } cmd.Flags().StringVarP(&loginToken, "token", "", "", "token provided to auto login") 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") + _ = 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 +162,18 @@ 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") } + // Browser/token login is an explicit "log me in as this user" action; suppress + // BREV_ACCESS_KEY so the freshly-saved JWT (not a stale env key) authenticates + // the post-login calls (GetCurrentUser, org selection, breadcrumbs). The + // --api-key path handles this via ResolveOrg promoting the flag key. + os.Unsetenv(auth.AccessKeyEnvVar) + tokens, _ := o.LoginStore.GetAuthTokens() if authProviderFlag != "" && authProviderFlag != "nvidia" && authProviderFlag != "legacy" { @@ -208,25 +216,26 @@ 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) + + // Activate the flag key so the org-resolution call (ListOrganizations) and the + // durable save both authenticate with it, not a pre-existing BREV_ACCESS_KEY. + os.Setenv(auth.AccessKeyEnvVar, apiKey) + org, err := register.ResolveOrgForAccessKey(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 } diff --git a/pkg/cmd/login/login_test.go b/pkg/cmd/login/login_test.go index 46fcec71..573e316e 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,70 @@ 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) +} + +// 'BREV_ACCESS_KEY=key-A brev login --api-key key-B' must resolve key-B's org: the +// flag key is promoted to the env var so the ListOrganizations call authenticates +// with it, not the pre-existing env key. +func TestRunLoginWithAPIKey_FlagKeyActivatesOverEnvKey(t *testing.T) { + t.Setenv(authpkg.AccessKeyEnvVar, 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.AccessKeyEnvVar) + 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") +} + +// Browser/token login is an explicit "log me in as this user" action and must +// suppress BREV_ACCESS_KEY for the whole transaction; otherwise post-login calls +// (org selection, breadcrumbs) authenticate with the env key, not the saved JWT. +func TestRunLogin_TokenLoginSuppressesEnvAccessKey(t *testing.T) { + t.Setenv(authpkg.AccessKeyEnvVar, 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.AccessKeyEnvVar), "browser/token login must clear BREV_ACCESS_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_test.go b/pkg/cmd/ls/ls_test.go index 0574e4e0..94ca1f9d 100644 --- a/pkg/cmd/ls/ls_test.go +++ b/pkg/cmd/ls/ls_test.go @@ -214,7 +214,7 @@ func TestRunLs_APIKeyRequiresCredentialOrg(t *testing.T) { if err == nil { t.Fatal("expected missing API key org error, got nil") } - if !strings.Contains(err.Error(), "api key auth requires an org id") { + if !strings.Contains(err.Error(), "org id missing") { t.Fatalf("expected API key org validation error, got %v", err) } if s.workspaceOrgID != "" { diff --git a/pkg/cmd/register/device_registration_store.go b/pkg/cmd/register/device_registration_store.go index 315dfb99..0838b098 100644 --- a/pkg/cmd/register/device_registration_store.go +++ b/pkg/cmd/register/device_registration_store.go @@ -20,31 +20,35 @@ const ( globalRegistrationDir = "/etc/brev" ) +const ( + RegistrationStatusPending = "pending" // used for retries + RegistrationStatusRegistered = "registered" +) + // DeviceRegistration is the persistent identity file for a registered device. // Fields align with the AddNodeResponse from dev-plane. type DeviceRegistration struct { - ExternalNodeID string `json:"external_node_id"` - DisplayName string `json:"display_name"` - OrgID string `json:"org_id"` - OrgName string `json:"org_name"` - DeviceID string `json:"device_id"` - RegisteredAt string `json:"registered_at"` - HardwareProfile HardwareProfile `json:"hardware_profile"` + ExternalNodeID string `json:"external_node_id"` + DisplayName string `json:"display_name"` + OrgID string `json:"org_id"` + OrgName string `json:"org_name"` + DeviceID string `json:"device_id"` + RegisteredAt string `json:"registered_at"` + HardwareProfile HardwareProfile `json:"hardware_profile"` + RegistrationToken string `json:"registration_token,omitempty"` + Status string `json:"status,omitempty"` } // RegistrationStore defines the contract for persisting device registration data. type RegistrationStore interface { Save(reg *DeviceRegistration) error - Load() (*DeviceRegistration, error) + Load(includeAll bool) (*DeviceRegistration, error) Delete() error Exists() (bool, error) } -// FileRegistrationStore implements RegistrationStore using the global /etc/brev/ path. type FileRegistrationStore struct{} -// NewFileRegistrationStore returns a FileRegistrationStore that reads/writes -// from /etc/brev/device_registration.json. func NewFileRegistrationStore() *FileRegistrationStore { return &FileRegistrationStore{} } @@ -72,8 +76,7 @@ func (s *FileRegistrationStore) Save(reg *DeviceRegistration) error { return sudoWriteFile(path, data) } -// Load reads the registration file and returns the parsed DeviceRegistration -func (s *FileRegistrationStore) Load() (*DeviceRegistration, error) { +func (s *FileRegistrationStore) Load(includeAll bool) (*DeviceRegistration, error) { path := s.path() exists, err := s.Exists() if !exists { @@ -86,7 +89,16 @@ func (s *FileRegistrationStore) Load() (*DeviceRegistration, error) { if err := files.ReadJSON(files.AppFs, path, ®); err != nil { return nil, breverrors.WrapAndTrace(err) } + if includeAll { + if reg.OrgID == "" && reg.DeviceID == "" { + return nil, breverrors.New("malformed registration") + } + return ®, nil + } if reg.ExternalNodeID == "" || reg.OrgID == "" { + if reg.Status == RegistrationStatusPending { + return nil, breverrors.New("device registration is incomplete; re-run 'brev register' to finish") + } return nil, breverrors.New("malformed registration") } return ®, nil diff --git a/pkg/cmd/register/device_registration_store_test.go b/pkg/cmd/register/device_registration_store_test.go index 39d7b1a2..2e1b085f 100644 --- a/pkg/cmd/register/device_registration_store_test.go +++ b/pkg/cmd/register/device_registration_store_test.go @@ -1,6 +1,7 @@ package register import ( + "strings" "testing" "github.com/brevdev/brev-cli/pkg/files" @@ -42,7 +43,7 @@ func Test_SaveAndLoadRegistration_RoundTrip(t *testing.T) { t.Fatalf("Save failed: %v", err) } - loaded, err := store.Load() + loaded, err := store.Load(false) if err != nil { t.Fatalf("Load failed: %v", err) } @@ -138,7 +139,7 @@ func Test_LoadRegistration_FailsWhenMissing(t *testing.T) { store := NewFileRegistrationStore() - _, err := store.Load() + _, err := store.Load(false) if err == nil { t.Error("expected error loading missing registration") } @@ -159,7 +160,7 @@ func Test_LoadRegistration_RejectsMissingExternalNodeID(t *testing.T) { t.Fatalf("Save failed: %v", err) } - _, err := store.Load() + _, err := store.Load(false) if err == nil { t.Fatal("expected error loading registration with empty ExternalNodeID") } @@ -180,7 +181,7 @@ func Test_LoadRegistration_RejectsMissingOrgID(t *testing.T) { t.Fatalf("Save failed: %v", err) } - _, err := store.Load() + _, err := store.Load(false) if err == nil { t.Fatal("expected error loading registration with empty OrgID") } @@ -197,3 +198,63 @@ func Test_DeleteRegistration_FailsWhenMissing(t *testing.T) { t.Error("expected error deleting missing registration") } } + +func Test_Load_IncludeAllReturnsPendingRecord(t *testing.T) { + cleanup := setupTestFs(t) + defer cleanup() + + store := NewFileRegistrationStore() + + pending := &DeviceRegistration{ + DisplayName: "My Spark", + OrgID: "org_xyz", + DeviceID: "device-uuid-123", + Status: RegistrationStatusPending, + } + if err := store.Save(pending); err != nil { + t.Fatalf("Save failed: %v", err) + } + + loaded, err := store.Load(true) + if err != nil { + t.Fatalf("Load failed: %v", err) + } + if loaded.DeviceID != "device-uuid-123" { + t.Errorf("DeviceID mismatch: got %s, want device-uuid-123", loaded.DeviceID) + } + if loaded.Status != RegistrationStatusPending { + t.Errorf("Status mismatch: got %q, want %q", loaded.Status, RegistrationStatusPending) + } + if loaded.ExternalNodeID != "" { + t.Errorf("pending record should have no ExternalNodeID, got %q", loaded.ExternalNodeID) + } + + if _, err := store.Load(false); err == nil { + t.Error("expected Load(false) to error on a pending record") + } +} + +func Test_Load_PendingRecordErrorMessage(t *testing.T) { + cleanup := setupTestFs(t) + defer cleanup() + + store := NewFileRegistrationStore() + + pending := &DeviceRegistration{ + DisplayName: "My Spark", + OrgID: "org_xyz", + DeviceID: "device-uuid-123", + Status: RegistrationStatusPending, + } + if err := store.Save(pending); err != nil { + t.Fatalf("Save failed: %v", err) + } + + _, err := store.Load(false) + if err == nil { + t.Fatal("expected Load() to error on a pending record") + } + if !strings.Contains(err.Error(), "incomplete") { + t.Errorf("expected 'incomplete' in error, got: %v", err) + } +} diff --git a/pkg/cmd/register/register.go b/pkg/cmd/register/register.go index 2ad88b43..092042cf 100644 --- a/pkg/cmd/register/register.go +++ b/pkg/cmd/register/register.go @@ -5,7 +5,7 @@ import ( "context" "errors" "fmt" - "os/user" + "os" "strings" "time" @@ -13,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" @@ -78,28 +79,48 @@ func defaultRegisterDeps() registerDeps { } } +type OrgLister interface { + ListOrganizations() ([]entity.Organization, error) +} + +func resolveAccessKey() string { + return strings.TrimSpace(os.Getenv(auth.AccessKeyEnvVar)) +} + var ( registerLong = `Register your device with NVIDIA Brev -This command sets up network connectivity and registers this machine with Brev. +This command registers this machine with Brev and brings up the Brev tunnel. +Registration no longer enables SSH; run 'brev enable-ssh' afterwards if you +want to SSH to this device. Two modes are supported: - • Interactive (default): run 'brev register' with no flags and follow prompts for device name, org, and options. - • Non-interactive: use any of --name, --org, or --ssh-port. No prompts; --name and --org are required. Use for scripts/CI.` + • Interactive (default): run 'brev register' with no flags and follow prompts for device name and org. + • 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_ACCESS_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 - # Non-interactive (any flag implies no prompts; --name and --org required) + # Non-interactive (--name and --org required) brev register --name my-node --org my-org - brev register --name my-node --org my-org --ssh-port 22` + + # Enable SSH access to this device after registering + brev enable-ssh` ) func NewCmdRegister(t *terminal.Terminal, store RegisterStore) *cobra.Command { var orgFlag string var nameFlag string - var sshPort int + var sshPort int // deprecated; accepted for backwards compatibility, no longer acted on var approveFlag bool + var registrationToken string cmd := &cobra.Command{ Annotations: map[string]string{"configuration": ""}, @@ -112,13 +133,14 @@ func NewCmdRegister(t *terminal.Terminal, store RegisterStore) *cobra.Command { RunE: func(cmd *cobra.Command, args []string) error { interactive := nameFlag == "" && orgFlag == "" && sshPort == 0 opts := registerOpts{ - interactive: interactive, - name: nameFlag, - orgName: orgFlag, - sshPort: int32(sshPort), - skipConfirm: approveFlag, + interactive: interactive, + name: nameFlag, + orgName: orgFlag, + skipConfirm: approveFlag, + registrationToken: registrationToken, } - return runRegister(cmd.Context(), t, store, opts, defaultRegisterDeps()) + deps := defaultRegisterDeps() + return runRegister(cmd.Context(), t, store, opts, deps) }, } @@ -126,22 +148,22 @@ func NewCmdRegister(t *terminal.Terminal, store RegisterStore) *cobra.Command { cmd.Flags().StringVarP(&nameFlag, "name", "n", "", "device name (required when using non-interactive mode)") cmd.Flags().IntVarP(&sshPort, "ssh-port", "p", 0, "SSH port (if ssh access is desired)") cmd.Flags().BoolVar(&approveFlag, "approve", false, "skip all confirmation prompts (assume yes)") + cmd.Flags().StringVar(®istrationToken, "registration-token", "", "Brev registration token") + _ = cmd.Flags().MarkDeprecated("ssh-port", "use 'brev enable-ssh' after registration to enable SSH access") return cmd } -// registerOpts carries mode and inputs: when interactive, name/orgName/sshPort are from prompts; otherwise from flags. +// registerOpts carries mode and inputs: when interactive, name/orgName are from prompts; otherwise from flags. type registerOpts struct { - interactive bool - name string - orgName string - sshPort int32 - skipConfirm bool + interactive bool + name string + orgName string + skipConfirm bool + registrationToken string } -// runRegister runs a single registration flow; the only difference by mode is whether we prompt or use opts. func runRegister(ctx context.Context, t *terminal.Terminal, s RegisterStore, opts registerOpts, deps registerDeps) error { //nolint:gocognit,gocyclo,funlen // ok - // Basic validation if !deps.platform.IsCompatible() { return breverrors.New("brev register is only supported on Linux") } @@ -149,28 +171,63 @@ func runRegister(ctx context.Context, t *terminal.Terminal, s RegisterStore, opt if err := deps.gater.Gate(t, deps.prompter, "Device registration", !opts.interactive || opts.skipConfirm); err != nil { return fmt.Errorf("sudo issue: %w", err) } + + accessKey := resolveAccessKey() 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 == "" && accessKey == "" { + return fmt.Errorf("in non-interactive mode --org is required unless --api-key is supplied") + } + } + if accessKey != "" { + if !auth.IsBrevAPIKey(accessKey) { + return breverrors.NewValidationError(fmt.Sprintf("api key must be a Brev API key (expected %s prefix); see 'brev login --api-key'", auth.BrevAPIKeyPrefix)) } + t.Vprintf(" %s\n", t.Green("Authenticating with API key.")) } - // Run through the login flow - brevUser, err := s.GetCurrentUser() - if err != nil { + // Verify the user is authenticated before performing any local side effects. + if _, err := s.GetCurrentUser(); err != nil { return breverrors.WrapAndTrace(err) } - // Check if the device is already registered - alreadyRegistered, err := deps.registrationStore.Exists() + var intendedOrg *entity.Organization + switch { + case accessKey != "": + o, err := ResolveOrgForAccessKey(s, opts.orgName) + if err != nil { + return err + } + intendedOrg = o + case !opts.interactive: + o, err := resolveOrg(s, opts.orgName) + if err != nil { + return err + } + intendedOrg = o + } + + // Check for an existing registration (confirmed or in-progress). + exists, err := deps.registrationStore.Exists() if err != nil { return breverrors.WrapAndTrace(err) } - if alreadyRegistered { + if exists { + reg, err := deps.registrationStore.Load(true) + if err != nil { + return breverrors.WrapAndTrace(err) + } + if intendedOrg != nil && intendedOrg.ID != reg.OrgID { + return orgMismatchError(reg, intendedOrg) + } + if reg.Status == RegistrationStatusPending { + return resumeRegistration(ctx, t, s, deps, reg) + } return checkExistingRegistration(ctx, t, s, deps) } - // Capture the device name var name string if opts.interactive { t.Vprint("") @@ -187,13 +244,13 @@ func runRegister(ctx context.Context, t *terminal.Terminal, s RegisterStore, opt return err //nolint:wrapcheck // do not present stack trace for this error } - // Capture the target organization + // Non-interactive already resolved intendedOrg above; interactive prompts. var org *entity.Organization - if opts.interactive { + if intendedOrg != nil { + org = intendedOrg + } else { t.Vprint("") org, err = resolveOrgInteractive(t, s, deps) - } else { - org, err = resolveOrg(s, opts.orgName) } if err != nil { return err @@ -226,48 +283,26 @@ func runRegister(ctx context.Context, t *terminal.Terminal, s RegisterStore, opt } } - // Perform the registration steps - reg, err := runRegisterSteps(ctx, t, s, name, org, deps) + // Generate the device ID here so a retry reuses it (AddNode is idempotent on device_id). + deviceID := uuid.New().String() + err = runRegisterSteps(ctx, t, s, name, org, deps, deviceID, opts.registrationToken) if err != nil { return err } - // Determine if SSH access should be enabled - enableSSH := false - sshPortForGrant := int32(0) - if opts.interactive { - enableSSH = deps.prompter.ConfirmYesNo("Would you like to enable SSH access to this device?") - if enableSSH { - sshPortForGrant = 0 // prompt for port - } - } else if opts.sshPort != 0 { - enableSSH = true - sshPortForGrant = opts.sshPort - } - - // Grant SSH access if requested - if enableSSH { - osUser, err := user.Current() - if err != nil { - return fmt.Errorf("failed to determine current Linux user: %w", err) - } - if err := grantSSHAccessWithPort(ctx, t, deps, s, reg, brevUser, osUser, sshPortForGrant, opts.interactive, opts.skipConfirm); err != nil { - t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: %v", err))) - } - } - + suggestEnableSSH(t) return nil } -// runRegisterSteps performs netbird install, hardware profile, AddNode, save registration, and runSetup. -// It does not prompt or enable SSH. Used by both flag-driven and prompt-driven flows. -func runRegisterSteps(ctx context.Context, t *terminal.Terminal, s RegisterStore, name string, org *entity.Organization, deps registerDeps) (*DeviceRegistration, error) { +// runRegisterSteps runs tunnel install, hardware profile, AddNode, persist, and +// setup. The pending write and deviceID-reuse rationale is at the pending site. +func runRegisterSteps(ctx context.Context, t *terminal.Terminal, s RegisterStore, name string, org *entity.Organization, deps registerDeps, deviceID, registrationToken string) error { t.Vprint("") t.Vprint(t.Yellow("[Step 1/5] Downloading and installing Brev tunnel...")) err := deps.netbird.Install() if err != nil { - return nil, fmt.Errorf("brev tunnel setup failed: %w", err) + return fmt.Errorf("brev tunnel setup failed: %w", err) } t.Vprintf("%s Brev tunnel ready.\n", t.Green(" ✓")) @@ -275,7 +310,7 @@ func runRegisterSteps(ctx context.Context, t *terminal.Terminal, s RegisterStore t.Vprint(t.Yellow("[Step 2/5] Collecting hardware profile...")) hwProfile, err := deps.hardwareProfiler.Profile() if err != nil { - return nil, fmt.Errorf("failed to collect hardware profile: %w", err) + return fmt.Errorf("failed to collect hardware profile: %w", err) } t.Vprintf("%s Hardware profile collected.\n", t.Green(" ✓")) t.Vprint("") @@ -284,7 +319,25 @@ func runRegisterSteps(ctx context.Context, t *terminal.Terminal, s RegisterStore t.Vprint("") t.Vprint(t.Yellow("[Step 3/5] Registering device with Brev...")) - deviceID := uuid.New().String() + + // A pending record written before AddNode (see resumeRegistration) makes the + // flow resumable: a crash/timeout after the node exists is retried with the + // same device ID, and AddNode is idempotent on device_id. + pending := &DeviceRegistration{ + DisplayName: name, + OrgID: org.ID, + OrgName: org.Name, + DeviceID: deviceID, + RegistrationToken: registrationToken, + HardwareProfile: *hwProfile, + Status: RegistrationStatusPending, + RegisteredAt: time.Now().UTC().Format(time.RFC3339), + // TODO use registration-token when backend API allows it + } + if err := deps.registrationStore.Save(pending); err != nil { + return fmt.Errorf("failed to write pending registration: %w", err) + } + client := deps.nodeClients.NewNodeClient(s, config.GlobalConfig.GetBrevPublicAPIURL()) addResp, err := client.AddNode(ctx, connect.NewRequest(&nodev1.AddNodeRequest{ OrganizationId: org.ID, @@ -293,30 +346,32 @@ func runRegisterSteps(ctx context.Context, t *terminal.Terminal, s RegisterStore NodeSpec: toProtoNodeSpec(hwProfile), })) if err != nil { - // dev-plane returns CodeAlreadyExists for a duplicate node name; surface - // its message directly, which already reads as "node already exists". var connectErr *connect.Error if errors.As(err, &connectErr) && connectErr.Code() == connect.CodeAlreadyExists { - return nil, errors.New(connectErr.Message()) + // delete pending registration to prevent stale dupe name + _ = deps.registrationStore.Delete() + return errors.New(connectErr.Message()) } - return nil, fmt.Errorf("failed to register node: %w", err) + return fmt.Errorf("failed to register node: %w", err) } node := addResp.Msg.GetExternalNode() reg := &DeviceRegistration{ - ExternalNodeID: node.GetExternalNodeId(), - DisplayName: name, - OrgID: org.ID, - OrgName: org.Name, - DeviceID: deviceID, - RegisteredAt: time.Now().UTC().Format(time.RFC3339), - HardwareProfile: *hwProfile, + ExternalNodeID: node.GetExternalNodeId(), + DisplayName: name, + OrgID: org.ID, + OrgName: org.Name, + DeviceID: deviceID, + RegistrationToken: registrationToken, + RegisteredAt: time.Now().UTC().Format(time.RFC3339), + HardwareProfile: *hwProfile, + Status: RegistrationStatusRegistered, } t.Vprint("") t.Vprint(t.Yellow("[Step 4/5] Storing registration data...")) if err := deps.registrationStore.Save(reg); err != nil { - return nil, fmt.Errorf("node registered but failed to save locally: %w", err) + return fmt.Errorf("node registered but failed to save locally: %w", err) } t.Vprint("") @@ -325,7 +380,7 @@ func runRegisterSteps(ctx context.Context, t *terminal.Terminal, s RegisterStore t.Vprintf("%s Node registered.\n", t.Green(" ✓")) t.Vprintf("%s Registration complete.\n", t.Green(" ✓")) - return reg, nil + return nil } func resolveOrgInteractive(t *terminal.Terminal, s RegisterStore, deps registerDeps) (*entity.Organization, error) { @@ -348,12 +403,40 @@ func resolveOrg(s RegisterStore, orgName string) (*entity.Organization, error) { return org, nil } -// checkExistingRegistration verifies connectivity for an already-registered node. -// It calls GetNode to check the server-side NetworkMemberStatus and ensures the -// local netbird service is running, starting it if necessary. Returns nil if -// the node is healthy, or an error describing what's wrong. +func ResolveOrgForAccessKey(s OrgLister, orgName string) (*entity.Organization, error) { + orgs, err := s.ListOrganizations() + if err != nil { + return nil, breverrors.WrapAndTrace(err) + } + org, err := singleOrgForAccessKey(orgs) + if err != nil { + return nil, err + } + if orgName != "" && org.Name != orgName { + return nil, breverrors.NewValidationError(fmt.Sprintf("access key does not belong to organization %q", orgName)) + } + return org, nil +} + +func singleOrgForAccessKey(orgs []entity.Organization) (*entity.Organization, error) { + if len(orgs) == 0 || len(orgs) > 1 { + return nil, breverrors.New("access key invalid") + } + return &orgs[0], nil +} + +func orgMismatchError(reg *DeviceRegistration, intended *entity.Organization) error { + existing := "this device is already registered in org" + if reg.Status == RegistrationStatusPending { + existing = "an incomplete registration exists for org" + } + return breverrors.NewValidationError(fmt.Sprintf( + "%s %s (%s), not %s (%s); run 'brev deregister' first to register in a different org", + existing, reg.OrgName, reg.OrgID, intended.Name, intended.ID)) +} + func checkExistingRegistration(ctx context.Context, t *terminal.Terminal, s RegisterStore, deps registerDeps) error { - reg, loadErr := deps.registrationStore.Load() + reg, loadErr := deps.registrationStore.Load(false) if loadErr != nil { return fmt.Errorf("this machine is already registered but the registration file could not be read: %w", loadErr) } @@ -363,11 +446,9 @@ func checkExistingRegistration(ctx context.Context, t *terminal.Terminal, s Regi t.Vprint(" Checking connectivity...") t.Vprint("") - // Check server-side connectivity status via GetNode. client := deps.nodeClients.NewNodeClient(s, config.GlobalConfig.GetBrevPublicAPIURL()) resp, err := client.GetNode(ctx, connect.NewRequest(&nodev1.GetNodeRequest{ ExternalNodeId: reg.ExternalNodeID, - OrganizationId: reg.OrgID, })) if err != nil { t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: could not fetch node status: %v", err))) @@ -424,58 +505,32 @@ func runSetup(node *nodev1.ExternalNode, t *terminal.Terminal, deps registerDeps } } -// grantSSHAccessWithPort enables SSH: shows confirm table, uses port or prompts if port is 0, then allocates port and grants access. -func grantSSHAccessWithPort(ctx context.Context, t *terminal.Terminal, deps registerDeps, tokenProvider externalnode.TokenProvider, reg *DeviceRegistration, brevUser *entity.User, osUser *user.User, port int32, interactive bool, skipConfirm bool) error { - brevUserName := brevUser.Username - if brevUserName == "" { - brevUserName = brevUser.Email - } - if brevUserName == "" { - brevUserName = brevUser.ID - } - +// resumeRegistration reuses the pending record's device ID. AddNode is +// idempotent on device_id, so this recovers when AddNode succeeded backend-side +// but the CLI never confirmed the ExternalNodeID. +func resumeRegistration(ctx context.Context, t *terminal.Terminal, s RegisterStore, deps registerDeps, pending *DeviceRegistration) error { t.Vprint("") t.Vprint(t.White("══════════════════════════════════════════════════")) - t.Vprint(t.White(" Enabling SSH access on this device")) + t.Vprint(t.White(" Resuming incomplete registration")) t.Vprint(t.White("══════════════════════════════════════════════════")) t.Vprint("") - if interactive && !skipConfirm { - t.Vprint(t.Green(" Please confirm before continuing:")) - t.Vprint("") - } - t.Vprintf(" %s %s\n", t.Green(fmt.Sprintf("%-14s", "Device:")), t.BoldBlue(reg.DisplayName+" ("+reg.ExternalNodeID+")")) - t.Vprintf(" %s %s\n", t.Green(fmt.Sprintf("%-14s", "Organization:")), t.BoldBlue(reg.OrgName+" ("+reg.OrgID+")")) - t.Vprintf(" %s %s\n", t.Green(fmt.Sprintf("%-14s", "Brev user:")), t.BoldBlue(brevUserName+" ("+brevUser.ID+")")) - t.Vprintf(" %s %s\n", t.Green(fmt.Sprintf("%-14s", "Linux user:")), t.BoldBlue(osUser.Username)) - - var err error - if port == 0 { - t.Vprint("") - port, err = PromptSSHPort(t) - if err != nil { - return fmt.Errorf("invalid SSH port: %w", err) - } - } else { - t.Vprintf(" %s %s\n", t.Green(fmt.Sprintf("%-14s", "SSH port:")), t.BoldBlue(fmt.Sprintf("%d", port))) - } + t.Vprintf(" %s %s\n", t.Green(fmt.Sprintf("%-14s", "Device:")), t.BoldBlue(pending.DisplayName)) + t.Vprintf(" %s %s\n", t.Green(fmt.Sprintf("%-14s", "Organization:")), t.BoldBlue(pending.OrgName+" ("+pending.OrgID+")")) + t.Vprintf(" %s %s\n", t.Green(fmt.Sprintf("%-14s", "Device ID:")), t.BoldBlue(pending.DeviceID)) t.Vprint("") + t.Vprint(" A previous registration attempt did not finish. Resuming.") - return grantSSHAccess(ctx, t, deps, tokenProvider, reg, brevUser, osUser, port) -} - -func grantSSHAccess(ctx context.Context, t *terminal.Terminal, deps registerDeps, tokenProvider externalnode.TokenProvider, reg *DeviceRegistration, brevUser *entity.User, osUser *user.User, port int32) error { - brevPortID, err := OpenSSHPort(ctx, t, deps.nodeClients, tokenProvider, reg, port) + org := &entity.Organization{ID: pending.OrgID, Name: pending.OrgName} + err := runRegisterSteps(ctx, t, s, pending.DisplayName, org, deps, pending.DeviceID, pending.RegistrationToken) if err != nil { - return fmt.Errorf("allocate SSH port failed: %w", err) + return err } - err = SetupAndRegisterNodeSSHAccess(ctx, t, deps.nodeClients, tokenProvider, reg, brevUser, osUser.Username, brevPortID) - if err != nil { - return fmt.Errorf("grant SSH failed: %w", err) - } + suggestEnableSSH(t) + return nil +} +func suggestEnableSSH(t *terminal.Terminal) { t.Vprint("") - t.Vprint(t.Green(fmt.Sprintf("SSH access enabled. You can now SSH to this device via: brev shell %s", reg.DisplayName))) - t.Vprint("") - return nil + t.Vprintf(" %s\n", t.Green("To enable SSH access to this device, run: brev enable-ssh")) } diff --git a/pkg/cmd/register/register_test.go b/pkg/cmd/register/register_test.go index d98b1a92..f267a75d 100644 --- a/pkg/cmd/register/register_test.go +++ b/pkg/cmd/register/register_test.go @@ -2,6 +2,7 @@ package register import ( "context" + "errors" "fmt" "net/http/httptest" "strings" @@ -11,7 +12,9 @@ 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" + breverrors "github.com/brevdev/brev-cli/pkg/errors" "github.com/brevdev/brev-cli/pkg/externalnode" "github.com/brevdev/brev-cli/pkg/sudo" "github.com/brevdev/brev-cli/pkg/terminal" @@ -72,7 +75,7 @@ func (m *mockRegistrationStore) Save(reg *DeviceRegistration) error { return nil } -func (m *mockRegistrationStore) Load() (*DeviceRegistration, error) { +func (m *mockRegistrationStore) Load(bool) (*DeviceRegistration, error) { if m.reg == nil { return nil, fmt.Errorf("no registration") } @@ -184,15 +187,33 @@ func testRegisterDeps(t *testing.T, svc *fakeNodeService, regStore RegistrationS }, server } +// testRegisterStore returns the default mock store used by most register tests: +// a logged-in user, a single resolvable org (org_123/TestOrg), and a token. +// Tests that need different orgs can override .org / .orgs after calling this. +func testRegisterStore() *mockRegisterStore { + return &mockRegisterStore{ + user: &entity.User{ID: "user_1"}, + org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, + token: "tok", + } +} + +// testPendingReg returns a pending DeviceRegistration for the given org/deviceID, +// the state left when a previous attempt didn't finish. +func testPendingReg(orgID, orgName, deviceID string) *DeviceRegistration { + return &DeviceRegistration{ + DisplayName: "My Spark", + OrgID: orgID, + OrgName: orgName, + DeviceID: deviceID, + Status: RegistrationStatusPending, + } +} + func Test_runRegister_HappyPath(t *testing.T) { regStore := &mockRegistrationStore{} - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - - token: "tok", - } + store := testRegisterStore() svc := &fakeNodeService{ addNodeFn: func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { @@ -223,11 +244,8 @@ func Test_runRegister_HappyPath(t *testing.T) { deps.setupRunner = setupRunner - SetTestSSHPort(22) - defer ClearTestSSHPort() - term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg", sshPort: 22} + opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg"} err := runRegister(context.Background(), term, store, opts, deps) if err != nil { t.Fatalf("runRegister failed: %v", err) @@ -242,7 +260,7 @@ func Test_runRegister_HappyPath(t *testing.T) { t.Fatal("expected registration to exist after successful register") } - reg, err := regStore.Load() + reg, err := regStore.Load(false) if err != nil { t.Fatalf("Load failed: %v", err) } @@ -272,11 +290,7 @@ func (f gaterFromFunc) Gate(t *terminal.Terminal, c terminal.Confirmer, reason s func Test_runRegister_UserCancels(t *testing.T) { // User cancel happens in interactive mode (sudo or confirm). Flag-driven has no prompts. regStore := &mockRegistrationStore{} - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - token: "tok", - } + store := testRegisterStore() svc := &fakeNodeService{} deps, server := testRegisterDeps(t, svc, regStore) defer server.Close() @@ -295,7 +309,7 @@ func Test_runRegister_UserCancels(t *testing.T) { }) term := terminal.New() - opts := registerOpts{interactive: true, name: "", orgName: "", sshPort: 0} + opts := registerOpts{interactive: true, name: "", orgName: ""} err := runRegister(context.Background(), term, store, opts, deps) if err == nil { t.Fatal("expected error when user declines sudo gate") @@ -369,12 +383,7 @@ func Test_runRegister_AlreadyRegistered(t *testing.T) { }, } - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - - token: "tok", - } + store := testRegisterStore() svc := &fakeNodeService{getNodeFn: tt.getNodeFn} deps, server := testRegisterDeps(t, svc, regStore) @@ -383,7 +392,7 @@ func Test_runRegister_AlreadyRegistered(t *testing.T) { term := terminal.New() // Pass the same name as the existing registration so we go through // the checkExistingRegistration path (not the different-name path). - opts := registerOpts{interactive: false, name: "Existing", orgName: "TestOrg", sshPort: 22} + opts := registerOpts{interactive: false, name: "Existing", orgName: "TestOrg"} err := runRegister(context.Background(), term, store, opts, deps) if err != nil { t.Fatalf("expected nil error, got: %v", err) @@ -397,28 +406,6 @@ func Test_runRegister_AlreadyRegistered(t *testing.T) { } } -func Test_runRegister_NoOrganization(t *testing.T) { - regStore := &mockRegistrationStore{} - - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: nil, - - token: "tok", - } - - svc := &fakeNodeService{} - deps, server := testRegisterDeps(t, svc, regStore) - defer server.Close() - - term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg", sshPort: 22} - err := runRegister(context.Background(), term, store, opts, deps) - if err == nil { - t.Fatal("expected error when no org exists") - } -} - func Test_runRegister_WithOrgFlag(t *testing.T) { regStore := &mockRegistrationStore{} @@ -451,11 +438,8 @@ func Test_runRegister_WithOrgFlag(t *testing.T) { defer server.Close() deps.setupRunner = setupRunner - SetTestSSHPort(22) - defer ClearTestSSHPort() - term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "SpecificOrg", sshPort: 22} + opts := registerOpts{interactive: false, name: "my-spark", orgName: "SpecificOrg"} err := runRegister(context.Background(), term, store, opts, deps) if err != nil { t.Fatalf("runRegister with --org failed: %v", err) @@ -465,7 +449,7 @@ func Test_runRegister_WithOrgFlag(t *testing.T) { t.Errorf("expected org_456, got %s", capturedOrgID) } - reg, err := regStore.Load() + reg, err := regStore.Load(false) if err != nil { t.Fatalf("Load failed: %v", err) } @@ -489,7 +473,7 @@ func Test_runRegister_WithOrgFlag_NotFound(t *testing.T) { defer server.Close() term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "NonexistentOrg", sshPort: 22} + opts := registerOpts{interactive: false, name: "my-spark", orgName: "NonexistentOrg"} err := runRegister(context.Background(), term, store, opts, deps) if err == nil { t.Fatal("expected error when org not found") @@ -499,51 +483,65 @@ func Test_runRegister_WithOrgFlag_NotFound(t *testing.T) { } } -func Test_runRegister_AddNodeFails(t *testing.T) { - regStore := &mockRegistrationStore{} - - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - - token: "tok", - } - - svc := &fakeNodeService{ - addNodeFn: func(_ *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { - return nil, connect.NewError(connect.CodeInternal, nil) - }, +func Test_runRegister_AddNodeFailure(t *testing.T) { + tests := []struct { + name string + code connect.Code + errMsg string + wantPending bool // true: record stays for resume; false: record cleared + wantErr string + }{ + {"Internal_StaysPending", connect.CodeInternal, "", true, ""}, + {"AlreadyExists_ClearsPending", connect.CodeAlreadyExists, "node with name my-spark already exists", false, "already exists"}, } + 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() + deps, server := testRegisterDeps(t, svc, regStore) + defer server.Close() - term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg", sshPort: 22} - err := runRegister(context.Background(), term, store, opts, deps) - if err == nil { - t.Fatal("expected error when AddNode fails") - } + term := terminal.New() + opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg"} + err := runRegister(context.Background(), term, store, opts, deps) + if err == nil { + t.Fatal("expected error on AddNode failure") + } + if tt.wantErr != "" && !strings.Contains(err.Error(), tt.wantErr) { + t.Errorf("expected error containing %q, got: %v", tt.wantErr, err) + } - // Registration should not exist on failure - exists, err := regStore.Exists() - if err != nil { - t.Fatalf("Exists error: %v", err) - } - if exists { - t.Error("registration should not exist after AddNode failure") + exists, _ := regStore.Exists() + if tt.wantPending != exists { + t.Errorf("wantPending=%v but exists=%v", tt.wantPending, exists) + } + if !tt.wantPending { + return + } + reg, loadErr := regStore.Load(true) + if loadErr != nil { + t.Fatalf("Load failed: %v", loadErr) + } + if reg.Status != RegistrationStatusPending { + t.Errorf("expected pending status, got %q", reg.Status) + } + if reg.DeviceID == "" { + t.Error("expected pending record to carry a device ID for retry") + } + }) } } func Test_runRegister_NoSetupCommand(t *testing.T) { regStore := &mockRegistrationStore{} - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - - token: "tok", - } + store := testRegisterStore() svc := &fakeNodeService{ addNodeFn: func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { @@ -566,11 +564,8 @@ func Test_runRegister_NoSetupCommand(t *testing.T) { deps.setupRunner = setupRunner - SetTestSSHPort(22) - defer ClearTestSSHPort() - term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg", sshPort: 22} + opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg"} err := runRegister(context.Background(), term, store, opts, deps) if err != nil { t.Fatalf("runRegister failed: %v", err) @@ -662,22 +657,79 @@ Peers count: 0/0 Connected` } } -func Test_runRegister_GrantSSH_retries_on_connection_error_then_succeeds(t *testing.T) { +func Test_runRegister_StepFailure(t *testing.T) { + tests := []struct { + name string + mutate func(*registerDeps) + errSubstr string + }{ + {"PlatformIncompatible", func(d *registerDeps) { d.platform = mockPlatform{compatible: false} }, "only supported on Linux"}, + {"HardwareProfilerFailure", func(d *registerDeps) { d.hardwareProfiler = &mockHardwareProfiler{err: fmt.Errorf("nvml init failed")} }, "hardware profile"}, + {"NetBirdInstallFailure", func(d *registerDeps) { d.netbird = mockNetBirdManager{err: fmt.Errorf("install failed")} }, "tunnel setup failed"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + deps, server := testRegisterDeps(t, &fakeNodeService{}, &mockRegistrationStore{}) + defer server.Close() + tt.mutate(&deps) + + opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg"} + err := runRegister(context.Background(), terminal.New(), testRegisterStore(), opts, deps) + if err == nil { + t.Fatal("expected error") + } + if !strings.Contains(err.Error(), tt.errSubstr) { + t.Errorf("expected error containing %q, got: %v", tt.errSubstr, err) + } + }) + } +} + +func Test_runRegister_NoNameNotRegistered(t *testing.T) { + // In flag-driven mode, missing --name and --org must error (no prompts). regStore := &mockRegistrationStore{} - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - token: "tok", + store := testRegisterStore() + + svc := &fakeNodeService{} + deps, server := testRegisterDeps(t, svc, regStore) + defer server.Close() + + term := terminal.New() + opts := registerOpts{interactive: false, name: "", orgName: ""} + err := runRegister(context.Background(), term, store, opts, deps) + if err == nil { + t.Fatal("expected error when no name/org in non-interactive mode") + } + if !strings.Contains(err.Error(), "non-interactive") || !strings.Contains(err.Error(), "--name") { + t.Errorf("expected non-interactive/--name error, got: %v", err) + } +} + +func Test_runRegister_ResumesPendingRegistration(t *testing.T) { + const pendingDeviceID = "device-uuid-pending" + pending := &DeviceRegistration{ + DisplayName: "My Spark", + OrgID: "org_123", + OrgName: "TestOrg", + DeviceID: pendingDeviceID, + RegistrationToken: "ui-token-pending", + Status: RegistrationStatusPending, } + regStore := &mockRegistrationStore{reg: pending} - var grantCalls int + store := testRegisterStore() + + var addNodeDeviceIDs []string + var addNodeCalls int svc := &fakeNodeService{ addNodeFn: func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { + addNodeCalls++ + addNodeDeviceIDs = append(addNodeDeviceIDs, req.GetDeviceId()) return &nodev1.AddNodeResponse{ ExternalNode: &nodev1.ExternalNode{ ExternalNodeId: "unode_abc", - OrganizationId: "org_123", + OrganizationId: req.GetOrganizationId(), Name: req.GetName(), DeviceId: req.GetDeviceId(), ConnectivityInfo: &nodev1.ConnectivityInfo{ @@ -686,270 +738,172 @@ func Test_runRegister_GrantSSH_retries_on_connection_error_then_succeeds(t *test }, }, nil }, - grantNodeSSHAccessFn: func(_ *nodev1.GrantNodeSSHAccessRequest) (*nodev1.GrantNodeSSHAccessResponse, error) { - grantCalls++ - if grantCalls < 2 { - return nil, connect.NewError(connect.CodeInternal, nil) - } - return &nodev1.GrantNodeSSHAccessResponse{}, nil - }, } deps, server := testRegisterDeps(t, svc, regStore) defer server.Close() - deps.prompter = mockConfirmer{confirm: true} - - SetTestSSHPort(22) - defer ClearTestSSHPort() - term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg", sshPort: 22} - err := runRegister(context.Background(), term, store, opts, deps) + // Use interactive mode so non-interactive --name/--org validation doesn't run + // before the Exists check; on resume the pending record's values are used. + err := runRegister(context.Background(), term, store, registerOpts{interactive: true}, deps) if err != nil { t.Fatalf("runRegister failed: %v", err) } - if grantCalls != 2 { - t.Errorf("expected GrantNodeSSHAccess to be called 2 times (retry once), got %d", grantCalls) + if addNodeCalls != 1 { + t.Fatalf("expected AddNode to be called once, got %d", addNodeCalls) + } + if len(addNodeDeviceIDs) != 1 || addNodeDeviceIDs[0] != pendingDeviceID { + t.Errorf("expected AddNode to reuse device ID %q, got %v", pendingDeviceID, addNodeDeviceIDs) } -} - -func Test_runRegister_GrantSSH_no_retry_on_permanent_error(t *testing.T) { - regStore := &mockRegistrationStore{} - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - token: "tok", + reg, loadErr := regStore.Load(true) + if loadErr != nil { + t.Fatalf("Load failed: %v", loadErr) + } + if reg.Status != RegistrationStatusRegistered { + t.Errorf("expected status %q after resume, got %q", RegistrationStatusRegistered, reg.Status) + } + if reg.ExternalNodeID != "unode_abc" { + t.Errorf("expected ExternalNodeID unode_abc, got %q", reg.ExternalNodeID) + } + if reg.DeviceID != pendingDeviceID { + t.Errorf("expected device ID to remain %q, got %q", pendingDeviceID, reg.DeviceID) } + if reg.RegistrationToken != "ui-token-pending" { + t.Errorf("expected registration token to be preserved, got %q", reg.RegistrationToken) + } +} - var grantCalls int +func Test_runRegister_PersistsRegistrationToken(t *testing.T) { + regStore := &mockRegistrationStore{} + store := testRegisterStore() svc := &fakeNodeService{ addNodeFn: func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { return &nodev1.AddNodeResponse{ ExternalNode: &nodev1.ExternalNode{ ExternalNodeId: "unode_abc", - OrganizationId: "org_123", + OrganizationId: req.GetOrganizationId(), Name: req.GetName(), DeviceId: req.GetDeviceId(), - ConnectivityInfo: &nodev1.ConnectivityInfo{ - RegistrationCommand: "netbird up --key abc", - }, }, }, nil }, - grantNodeSSHAccessFn: func(_ *nodev1.GrantNodeSSHAccessRequest) (*nodev1.GrantNodeSSHAccessResponse, error) { - grantCalls++ - return nil, connect.NewError(connect.CodePermissionDenied, nil) - }, } deps, server := testRegisterDeps(t, svc, regStore) defer server.Close() - deps.prompter = mockConfirmer{confirm: true} - - SetTestSSHPort(22) - defer ClearTestSSHPort() - term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg", sshPort: 22} + opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg", registrationToken: "ui-token-xyz"} err := runRegister(context.Background(), term, store, opts, deps) if err != nil { - t.Fatalf("runRegister should not fail the overall flow when SSH grant fails: %v", err) + t.Fatalf("runRegister failed: %v", err) } - if grantCalls != 1 { - t.Errorf("expected GrantNodeSSHAccess to be called once (no retry on permanent error), got %d", grantCalls) + reg, loadErr := regStore.Load(true) + if loadErr != nil { + t.Fatalf("Load failed: %v", loadErr) } -} - -func Test_runRegister_NameValidation(t *testing.T) { - tests := []struct { - name string - input string - wantErr bool - errSubstr string - }{ - {"Valid", "my-dgx-spark", false, ""}, - {"WithDots", "node.local.1", false, ""}, - {"WithUnderscore", "my_node", false, ""}, - {"Spaces", "My Spark", true, "letters, digits"}, - {"ShellInjection", "$(whoami)", true, "letters, digits"}, - {"PathTraversal", "../etc/passwd", true, "letters, digits"}, - {"Backticks", "`rm -rf`", true, "letters, digits"}, - {"Semicolon", "a;rm -rf /", true, "letters, digits"}, - {"LeadingHyphen", "-node", true, "start with"}, - {"LeadingDot", ".hidden", true, "start with"}, - {"TooLong", strings.Repeat("a", 64), true, "63 characters"}, - {"Empty", "", true, "--name"}, // flag-driven rejects empty name with this message + if reg.RegistrationToken != "ui-token-xyz" { + t.Errorf("expected RegistrationToken to be persisted, got %q", reg.RegistrationToken) } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - regStore := &mockRegistrationStore{} - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - token: "tok", - } - - svc := &fakeNodeService{ - addNodeFn: func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { - return &nodev1.AddNodeResponse{ - ExternalNode: &nodev1.ExternalNode{ - ExternalNodeId: "unode_abc", - OrganizationId: "org_123", - Name: req.GetName(), - DeviceId: req.GetDeviceId(), - }, - }, nil - }, - } - - deps, server := testRegisterDeps(t, svc, regStore) - defer server.Close() - - SetTestSSHPort(22) - defer ClearTestSSHPort() - - term := terminal.New() - var err error - opts := registerOpts{interactive: false, name: tt.input, orgName: "TestOrg", sshPort: 22} - err = runRegister(context.Background(), term, store, opts, deps) - if tt.wantErr { - if err == nil { - t.Fatal("expected error, got nil") - } - if !strings.Contains(err.Error(), tt.errSubstr) { - t.Errorf("expected error containing %q, got: %v", tt.errSubstr, err) - } - } else if err != nil { - t.Errorf("unexpected error: %v", err) - } - }) + if reg.Status != RegistrationStatusRegistered { + t.Errorf("expected registered status, got %q", reg.Status) } } -func Test_runRegister_PlatformIncompatible(t *testing.T) { - regStore := &mockRegistrationStore{} +const testAccessKey = auth.BrevAPIKeyPrefix + "test-key" - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - token: "tok", +func accessKeyAddNodeFn(t *testing.T) func(*nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { + t.Helper() + return func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { + if req.GetOrganizationId() != "org_123" { + t.Errorf("unexpected org: %s", req.GetOrganizationId()) + } + return &nodev1.AddNodeResponse{ + ExternalNode: &nodev1.ExternalNode{ + ExternalNodeId: "unode_abc", + OrganizationId: req.GetOrganizationId(), + Name: req.GetName(), + DeviceId: req.GetDeviceId(), + }, + }, nil } +} - svc := &fakeNodeService{} - deps, server := testRegisterDeps(t, svc, regStore) - defer server.Close() - - deps.platform = mockPlatform{compatible: false} - - term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg", sshPort: 22} - err := runRegister(context.Background(), term, store, opts, deps) - if err == nil { - t.Fatal("expected error when platform is incompatible") - } - if !strings.Contains(err.Error(), "only supported on Linux") { - t.Errorf("expected platform incompatibility error, got: %v", err) - } +func ensureNoAccessKeyEnv(t *testing.T) { + t.Helper() + t.Setenv(auth.AccessKeyEnvVar, "") } -func Test_runRegister_HardwareProfilerFailure(t *testing.T) { - regStore := &mockRegistrationStore{} +func Test_resolveAccessKey(t *testing.T) { + ensureNoAccessKeyEnv(t) - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - token: "tok", + if got := resolveAccessKey(); got != "" { + t.Errorf("expected empty when no env, got %q", got) } - svc := &fakeNodeService{} - deps, server := testRegisterDeps(t, svc, regStore) - defer server.Close() - - deps.hardwareProfiler = &mockHardwareProfiler{err: fmt.Errorf("nvml init failed")} - - term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg", sshPort: 22} - err := runRegister(context.Background(), term, store, opts, deps) - if err == nil { - t.Fatal("expected error when hardware profiler fails") - } - if !strings.Contains(err.Error(), "hardware profile") { - t.Errorf("expected hardware profile error, got: %v", err) + t.Setenv(auth.AccessKeyEnvVar, auth.BrevAPIKeyPrefix+"env-key") + if got := resolveAccessKey(); got != auth.BrevAPIKeyPrefix+"env-key" { + t.Errorf("expected env value, got %q", got) } } -func Test_runRegister_NetBirdInstallFailure(t *testing.T) { - regStore := &mockRegistrationStore{} +func Test_runRegister_AccessKeyEnv_Fallback(t *testing.T) { + ensureNoAccessKeyEnv(t) + t.Setenv(auth.AccessKeyEnvVar, testAccessKey) - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - token: "tok", - } + regStore := &mockRegistrationStore{} + store := testRegisterStore() + svc := &fakeNodeService{addNodeFn: accessKeyAddNodeFn(t)} - svc := &fakeNodeService{} deps, server := testRegisterDeps(t, svc, regStore) defer server.Close() - - deps.netbird = mockNetBirdManager{err: fmt.Errorf("install failed")} - term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg", sshPort: 22} - err := runRegister(context.Background(), term, store, opts, deps) - if err == nil { - t.Fatal("expected error when NetBird install fails") - } - if !strings.Contains(err.Error(), "tunnel setup failed") { - t.Errorf("expected tunnel setup error, got: %v", err) + opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg"} + if err := runRegister(context.Background(), term, store, opts, deps); err != nil { + t.Fatalf("runRegister failed: %v", err) } } -func Test_runRegister_NoNameNotRegistered(t *testing.T) { - // In flag-driven mode, missing --name and --org must error (no prompts). - regStore := &mockRegistrationStore{} +func Test_runRegister_AccessKeyInvalid(t *testing.T) { + t.Setenv(auth.AccessKeyEnvVar, "not-a-brev-key") - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - token: "tok", - } + regStore := &mockRegistrationStore{} + store := testRegisterStore() + svc := &fakeNodeService{addNodeFn: accessKeyAddNodeFn(t)} - svc := &fakeNodeService{} deps, server := testRegisterDeps(t, svc, regStore) defer server.Close() - term := terminal.New() - opts := registerOpts{interactive: false, name: "", orgName: "", sshPort: 22} + opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg"} err := runRegister(context.Background(), term, store, opts, deps) if err == nil { - t.Fatal("expected error when no name/org in non-interactive mode") + t.Fatal("expected error for invalid access key") } - if !strings.Contains(err.Error(), "non-interactive") || !strings.Contains(err.Error(), "--name") { - t.Errorf("expected non-interactive/--name error, got: %v", err) + var ve breverrors.ValidationError + if !errors.As(err, &ve) { + t.Errorf("expected a ValidationError, got %T: %v", err, err) + } + if !strings.Contains(err.Error(), auth.BrevAPIKeyPrefix) { + t.Errorf("expected error to mention the %s prefix, got: %v", auth.BrevAPIKeyPrefix, err) } } -func Test_runRegister_NoNameAlreadyRegistered(t *testing.T) { +func Test_runRegister_AccessKey_SeedsSessionOnAlreadyRegistered(t *testing.T) { + t.Setenv(auth.AccessKeyEnvVar, testAccessKey) + regStore := &mockRegistrationStore{ reg: &DeviceRegistration{ ExternalNodeID: "unode_existing", - DisplayName: "Existing Device", + DisplayName: "Existing", OrgID: "org_123", + Status: RegistrationStatusRegistered, }, } - - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - token: "tok", - } - + store := testRegisterStore() svc := &fakeNodeService{ getNodeFn: func(req *nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) { return &nodev1.GetNodeResponse{ @@ -965,167 +919,243 @@ func Test_runRegister_NoNameAlreadyRegistered(t *testing.T) { deps, server := testRegisterDeps(t, svc, regStore) defer server.Close() + term := terminal.New() + opts := registerOpts{interactive: false, name: "Existing", orgName: "TestOrg"} + if err := runRegister(context.Background(), term, store, opts, deps); err != nil { + t.Fatalf("runRegister failed: %v", err) + } +} + +// --- Access key org scoping --- + +func Test_resolveOrgForAccessKey(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, "", "", "access key invalid"}, + {"multiple, name matches one", []entity.Organization{{ID: "org_1", Name: "Alpha"}, {ID: "org_2", Name: "Beta"}}, "Beta", "", "access key invalid"}, + {"multiple, no name", []entity.Organization{{ID: "org_1", Name: "Alpha"}, {ID: "org_2", Name: "Beta"}}, "", "", "access key invalid"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + s := &mockRegisterStore{orgs: tt.orgs} + got, err := ResolveOrgForAccessKey(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) + } + }) + } +} + +func Test_runRegister_AccessKey_OrgMismatch(t *testing.T) { + t.Setenv(auth.AccessKeyEnvVar, testAccessKey) + regStore := &mockRegistrationStore{} + store := testRegisterStore() + svc := &fakeNodeService{addNodeFn: accessKeyAddNodeFn(t)} + + deps, server := testRegisterDeps(t, svc, regStore) + defer server.Close() term := terminal.New() - opts := registerOpts{interactive: false, name: "Existing", orgName: "TestOrg", sshPort: 22} + opts := registerOpts{interactive: false, name: "my-spark", orgName: "OtherOrg"} err := runRegister(context.Background(), term, store, opts, deps) - if err != nil { - t.Fatalf("expected nil error when already registered with no name, got: %v", err) + if err == nil { + t.Fatal("expected error when --org doesn't match the access key's org") } + if !strings.Contains(err.Error(), "does not belong to organization") { + t.Errorf("expected org mismatch error, got: %v", err) + } +} - // Registration should still exist - exists, _ := regStore.Exists() - if !exists { - t.Error("expected registration to still exist") +func Test_runRegister_AccessKey_NoOrgFlag_UsesKeyOrg(t *testing.T) { + t.Setenv(auth.AccessKeyEnvVar, testAccessKey) + + regStore := &mockRegistrationStore{} + store := testRegisterStore() + var gotOrgID string + svc := &fakeNodeService{addNodeFn: func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { + gotOrgID = req.GetOrganizationId() + return &nodev1.AddNodeResponse{ExternalNode: &nodev1.ExternalNode{ExternalNodeId: "unode_abc", OrganizationId: req.GetOrganizationId(), Name: req.GetName(), DeviceId: req.GetDeviceId()}}, nil + }} + + deps, server := testRegisterDeps(t, svc, regStore) + defer server.Close() + term := terminal.New() + // No --org; the key's org (org_123) should be used. + opts := registerOpts{interactive: false, name: "my-spark"} + if err := runRegister(context.Background(), term, store, opts, deps); err != nil { + t.Fatalf("runRegister failed: %v", err) + } + if gotOrgID != "org_123" { + t.Errorf("expected AddNode to use the key's org org_123, got %s", gotOrgID) } } -func Test_runRegister_OpenSSHPort(t *testing.T) { // nolint:funlen, gocyclo, gocognit // test +func Test_runRegister_OrgMismatch(t *testing.T) { tests := []struct { - name string - port int32 - openFn func(*nodev1.OpenPortRequest) (*nodev1.OpenPortResponse, error) - verify func(t *testing.T, openReq *nodev1.OpenPortRequest, grantReq *nodev1.GrantNodeSSHAccessRequest, reg *mockRegistrationStore, err error) + name string + status string + useAccessKey bool // set BREV_ACCESS_KEY for the new org + useOrgFlag bool // pass --org for the new org + wantWording string // pending -> "incomplete registration"; registered -> "already registered" }{ - { - name: "SendsCorrectArgs", - port: 2222, - openFn: func(req *nodev1.OpenPortRequest) (*nodev1.OpenPortResponse, error) { - return &nodev1.OpenPortResponse{ - Port: &nodev1.Port{ - PortId: "port_ssh", - Protocol: req.GetProtocol(), - PortNumber: req.GetPortNumber(), - }, - }, nil - }, - verify: func(t *testing.T, openReq *nodev1.OpenPortRequest, _ *nodev1.GrantNodeSSHAccessRequest, _ *mockRegistrationStore, err error) { - t.Helper() - if err != nil { - t.Fatalf("runRegister failed: %v", err) - } - if openReq == nil { - t.Fatal("expected OpenPort to be called") - } - if openReq.GetExternalNodeId() != "unode_abc" { - t.Errorf("expected node ID unode_abc, got %s", openReq.GetExternalNodeId()) - } - if openReq.GetProtocol() != nodev1.PortProtocol_PORT_PROTOCOL_TCP { - t.Errorf("expected PORT_PROTOCOL_TCP, got %s", openReq.GetProtocol()) - } - if openReq.GetPortNumber() != 2222 { - t.Errorf("expected port 2222, got %d", openReq.GetPortNumber()) - } - }, - }, - { - name: "FailureIsSoftError", - port: 22, - openFn: func(_ *nodev1.OpenPortRequest) (*nodev1.OpenPortResponse, error) { - return nil, connect.NewError(connect.CodeInternal, fmt.Errorf("skybridge unavailable")) - }, - verify: func(t *testing.T, _ *nodev1.OpenPortRequest, _ *nodev1.GrantNodeSSHAccessRequest, regStore *mockRegistrationStore, err error) { - t.Helper() - if err != nil { - t.Fatalf("registration should succeed even when OpenSSHPort fails (soft error), got: %v", err) - } - exists, _ := regStore.Exists() - if !exists { - t.Error("expected registration to still exist after OpenSSHPort failure") - } - }, - }, - { - name: "InvalidPortNoAPICall", - port: 99999, - verify: func(t *testing.T, openReq *nodev1.OpenPortRequest, _ *nodev1.GrantNodeSSHAccessRequest, regStore *mockRegistrationStore, err error) { - t.Helper() - if err != nil { - t.Fatalf("registration should succeed even when SSH port is invalid (soft error), got: %v", err) - } - if openReq != nil { - t.Error("expected OpenPort NOT to be called for invalid port") - } - exists, _ := regStore.Exists() - if !exists { - t.Error("expected registration to still exist after invalid port") - } - }, - }, - { - name: "GrantRequestHasNoPort", - port: 22, - verify: func(t *testing.T, _ *nodev1.OpenPortRequest, grantReq *nodev1.GrantNodeSSHAccessRequest, _ *mockRegistrationStore, err error) { - t.Helper() - if err != nil { - t.Fatalf("runRegister failed: %v", err) - } - if grantReq == nil { - t.Fatal("expected GrantNodeSSHAccess to be called") - } - if grantReq.GetExternalNodeId() != "unode_abc" { - t.Errorf("expected node ID unode_abc, got %s", grantReq.GetExternalNodeId()) - } - if grantReq.GetUserId() != "user_1" { - t.Errorf("expected user ID user_1, got %s", grantReq.GetUserId()) - } - }, - }, + {"AccessKey_Pending", RegistrationStatusPending, true, false, "incomplete registration"}, + {"OrgFlag_Pending", RegistrationStatusPending, false, true, "incomplete registration"}, + {"AccessKey_AlreadyRegistered", RegistrationStatusRegistered, true, false, "already registered"}, } - for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - regStore := &mockRegistrationStore{} + ensureNoAccessKeyEnv(t) + if tt.useAccessKey { + t.Setenv(auth.AccessKeyEnvVar, testAccessKey) + } + regStore := &mockRegistrationStore{reg: &DeviceRegistration{ + ExternalNodeID: "unode_existing", + DisplayName: "My Spark", + OrgID: "org_other", + OrgName: "OtherOrg", + DeviceID: "dev-pending", + Status: tt.status, + }} store := &mockRegisterStore{ user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, + orgs: []entity.Organization{{ID: "org_123", Name: "TestOrg"}}, // new org token: "tok", } - var gotOpenReq *nodev1.OpenPortRequest - var gotGrantReq *nodev1.GrantNodeSSHAccessRequest + var addNodeCalls int svc := &fakeNodeService{ - addNodeFn: func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { - 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", - }, - }, - }, nil - }, - openPortFn: func(req *nodev1.OpenPortRequest) (*nodev1.OpenPortResponse, error) { - gotOpenReq = req - if tt.openFn != nil { - return tt.openFn(req) - } - return &nodev1.OpenPortResponse{ - Port: &nodev1.Port{PortId: "port_ssh", Protocol: req.GetProtocol(), PortNumber: req.GetPortNumber()}, - }, nil - }, - grantNodeSSHAccessFn: func(req *nodev1.GrantNodeSSHAccessRequest) (*nodev1.GrantNodeSSHAccessResponse, error) { - gotGrantReq = req - return &nodev1.GrantNodeSSHAccessResponse{}, nil + addNodeFn: func(_ *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { + addNodeCalls++ + return nil, fmt.Errorf("AddNode should not be called on org mismatch") }, } deps, server := testRegisterDeps(t, svc, regStore) defer server.Close() - deps.prompter = mockConfirmer{confirm: true} - - SetTestSSHPort(tt.port) - defer ClearTestSSHPort() - term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "TestOrg", sshPort: tt.port} + opts := registerOpts{interactive: false, name: "My Spark", orgName: "TestOrg"} err := runRegister(context.Background(), term, store, opts, deps) - - tt.verify(t, gotOpenReq, gotGrantReq, regStore, err) + 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) + } + if addNodeCalls != 0 { + t.Errorf("AddNode must not be called on mismatch, got %d", addNodeCalls) + } }) } } + +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) { + addNodeCalls++ + return nil, connect.NewError(connect.CodeInternal, nil) + }, + } + + deps, server := testRegisterDeps(t, svc, regStore) + defer server.Close() + + term := terminal.New() + err := runRegister(context.Background(), term, store, registerOpts{interactive: true}, deps) + if err == nil { + t.Fatal("expected error when AddNode fails during resume") + } + if addNodeCalls != 1 { + t.Fatalf("expected AddNode called once, got %d", addNodeCalls) + } + + reg, loadErr := regStore.Load(true) + 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) + } + if reg.ExternalNodeID != "" { + t.Errorf("expected no ExternalNodeID after failed resume, got %q", reg.ExternalNodeID) + } +} + +func Test_runRegister_OrgFlag_PendingOrgMatch_Resumes(t *testing.T) { + const pendingDeviceID = "device-uuid-pending" + regStore := &mockRegistrationStore{reg: testPendingReg("org_123", "TestOrg", pendingDeviceID)} + store := &mockRegisterStore{ + user: &entity.User{ID: "user_1"}, + orgs: []entity.Organization{{ID: "org_123", Name: "TestOrg"}}, + token: "tok", + } + + var gotDeviceID string + svc := &fakeNodeService{ + addNodeFn: func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { + gotDeviceID = req.GetDeviceId() + return &nodev1.AddNodeResponse{ + ExternalNode: &nodev1.ExternalNode{ + ExternalNodeId: "unode_abc", + OrganizationId: req.GetOrganizationId(), + Name: req.GetName(), + DeviceId: req.GetDeviceId(), + }, + }, nil + }, + } + + deps, server := testRegisterDeps(t, svc, regStore) + defer server.Close() + + term := terminal.New() + opts := registerOpts{interactive: false, name: "My Spark", orgName: "TestOrg"} + if err := runRegister(context.Background(), term, store, opts, deps); err != nil { + t.Fatalf("runRegister failed: %v", err) + } + if gotDeviceID != pendingDeviceID { + t.Errorf("expected resume to reuse device ID %q, got %q", pendingDeviceID, gotDeviceID) + } + reg, loadErr := regStore.Load(true) + if loadErr != nil { + t.Fatalf("Load failed: %v", loadErr) + } + if reg.Status != RegistrationStatusRegistered { + t.Errorf("expected registered status after resume, got %q", reg.Status) + } +} diff --git a/pkg/cmd/register/sshkeys.go b/pkg/cmd/register/sshkeys.go index 1188766d..2ca9bb91 100644 --- a/pkg/cmd/register/sshkeys.go +++ b/pkg/cmd/register/sshkeys.go @@ -30,7 +30,7 @@ func SelectNodeFromList(ctx context.Context, t *terminal.Terminal, prompter term return nil, fmt.Errorf("no nodes found in organization") } var thisNodeID string - if reg, err := registrationStore.Load(); err == nil && reg != nil { + if reg, err := registrationStore.Load(false); err == nil && reg != nil { thisNodeID = reg.ExternalNodeID } t.Vprint("") diff --git a/pkg/cmd/revokessh/revokessh_test.go b/pkg/cmd/revokessh/revokessh_test.go index 447c53cb..71165928 100644 --- a/pkg/cmd/revokessh/revokessh_test.go +++ b/pkg/cmd/revokessh/revokessh_test.go @@ -43,7 +43,7 @@ func (m *mockRegistrationStore) Save(reg *register.DeviceRegistration) error { return nil } -func (m *mockRegistrationStore) Load() (*register.DeviceRegistration, error) { +func (m *mockRegistrationStore) Load(bool) (*register.DeviceRegistration, error) { if m.reg == nil { return nil, fmt.Errorf("no registration") } From cd860191e49bb8bf7be66d62e4acc1682c2e8504 Mon Sep 17 00:00:00 2001 From: Pratik Patel Date: Thu, 27 Aug 2026 08:06:56 -0700 Subject: [PATCH 2/2] Wire LinuxUser cache into enable-ssh and grant-ssh - enable-ssh: reuse cached Linux user so repeat grants stay consistent; cache the user after a successful grant - grant-ssh: interactive picker pre-selects the cached user with a confirm prompt; final choice is cached - register_test: fold PersistsRegistrationToken into HappyPath, drop redundant APIKeyEnv_Fallback (asserted less than NoOrgFlag_UsesKeyOrg) GetCachedLinuxUser/SaveCachedLinuxUser existed since BRE2-818 without any consumer; this wires them for the first time. --- pkg/cmd/enablessh/enablessh.go | 20 +++++++- pkg/cmd/grantssh/grantssh.go | 34 ++++++++++++-- pkg/cmd/grantssh/grantssh_test.go | 78 +++++++++++++++++++++++++++++++ 3 files changed, 126 insertions(+), 6 deletions(-) diff --git a/pkg/cmd/enablessh/enablessh.go b/pkg/cmd/enablessh/enablessh.go index 6b6ac548..504ad441 100644 --- a/pkg/cmd/enablessh/enablessh.go +++ b/pkg/cmd/enablessh/enablessh.go @@ -25,6 +25,8 @@ import ( type EnableSSHStore interface { GetCurrentUser() (*entity.User, error) GetAccessToken() (string, error) + GetCachedLinuxUser() (string, error) + SaveCachedLinuxUser(linuxUser string) error } // enableSSHDeps bundles the side-effecting dependencies of runEnableSSH so they @@ -76,7 +78,7 @@ func runEnableSSH(ctx context.Context, t *terminal.Terminal, s EnableSSHStore, d return breverrors.WrapAndTrace(err) } - return enableSSH(ctx, t, deps, s, reg, brevUser) + return enableSSH(ctx, t, s, deps, s, reg, brevUser) } // enableSSH grants SSH access to the given node for the current Brev user. @@ -84,19 +86,28 @@ func runEnableSSH(ctx context.Context, t *terminal.Terminal, s EnableSSHStore, d func enableSSH( ctx context.Context, t *terminal.Terminal, + s EnableSSHStore, deps enableSSHDeps, tokenProvider externalnode.TokenProvider, reg *register.DeviceRegistration, brevUser *entity.User, ) error { + // Reuse the Linux user cached by a previous enable-ssh on this machine so + // repeat grants stay consistent; fall back to the current OS user. + cachedLinuxUser, err := s.GetCachedLinuxUser() + if err != nil { + return breverrors.WrapAndTrace(err) + } linuxUser, err := user.Current() if err != nil { return fmt.Errorf("failed to determine current Linux user: %w", err) } linuxUsername := linuxUser.Username + if cachedLinuxUser != "" { + linuxUsername = cachedLinuxUser + } checkSSHDaemon(t) - t.Vprint("") t.Vprint(t.Green("Enabling SSH access on this device")) t.Vprint("") @@ -119,6 +130,11 @@ func enableSSH( return fmt.Errorf("enable SSH failed: %w", err) } + // Cache the Linux user so future grant-ssh prompts can pre-select it. + if err := s.SaveCachedLinuxUser(linuxUsername); err != nil { + t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: could not cache Linux user: %v", 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 } diff --git a/pkg/cmd/grantssh/grantssh.go b/pkg/cmd/grantssh/grantssh.go index 35985834..4f286399 100644 --- a/pkg/cmd/grantssh/grantssh.go +++ b/pkg/cmd/grantssh/grantssh.go @@ -32,6 +32,8 @@ type GrantSSHStore interface { GetAccessToken() (string, error) ListOrganizationMembers(ctx context.Context, orgID string) ([]*nodev1.OrganizationMember, error) GetUserByID(userID string) (*entity.User, error) + GetCachedLinuxUser() (string, error) + SaveCachedLinuxUser(linuxUser string) error } // grantSSHDeps bundles the side-effecting dependencies of runGrantSSH so they @@ -183,12 +185,13 @@ func runGrantSSH(ctx context.Context, t *terminal.Terminal, s GrantSSHStore, opt return err } linuxUserOptions := uniqueLinuxUsersFromNodeSSHAccess(node) - if len(linuxUserOptions) > 0 { - t.Vprint("") - linuxUser = deps.prompter.Select("Select Linux user on the node", linuxUserOptions) - } else { + 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)") } + linuxUser, err = selectCachedLinuxUser(t, deps.prompter, s, linuxUserOptions) + if err != nil { + return err + } } else { selectedUser, err = findUserByIDOrEmail(orgMembers, opts.userIDOrEmail) if err != nil { @@ -280,6 +283,29 @@ func uniqueLinuxUsersFromNodeSSHAccess(node *nodev1.ExternalNode) []string { return slices.Collect(maps.Keys(linuxUsers)) } +// selectCachedLinuxUser prompts for a Linux user from the node's options, +// pre-selecting the user cached by a previous enable-ssh on this machine, and +// caches the final choice so repeat grants default to it. +func selectCachedLinuxUser(t *terminal.Terminal, prompter terminal.Selector, s GrantSSHStore, options []string) (string, error) { + cached, err := s.GetCachedLinuxUser() + if err != nil { + return "", breverrors.WrapAndTrace(err) + } + if cached != "" && slices.Contains(options, cached) { + t.Vprint("") + confirm := prompter.Select(fmt.Sprintf("Linux user on the node [%s]", cached), []string{"Yes, use " + cached, "No, choose another"}) + if confirm == "Yes, use "+cached { + return cached, nil + } + } + t.Vprint("") + selected := prompter.Select("Select Linux user on the node", options) + if err := s.SaveCachedLinuxUser(selected); err != nil { + t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: could not cache Linux user: %v", err))) + } + return selected, nil +} + func findUserByIDOrEmail(members []resolvedMember, idOrEmail string) (*entity.User, error) { idOrEmail = strings.TrimSpace(strings.ToLower(idOrEmail)) for _, r := range members { diff --git a/pkg/cmd/grantssh/grantssh_test.go b/pkg/cmd/grantssh/grantssh_test.go index e1e6441a..40559368 100644 --- a/pkg/cmd/grantssh/grantssh_test.go +++ b/pkg/cmd/grantssh/grantssh_test.go @@ -14,6 +14,8 @@ import ( "github.com/brevdev/brev-cli/pkg/entity" "github.com/brevdev/brev-cli/pkg/externalnode" "github.com/brevdev/brev-cli/pkg/terminal" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) // mock types for grantSSHDeps interfaces @@ -68,6 +70,10 @@ type mockGrantSSHStore struct { members []*nodev1.OrganizationMember users map[string]*entity.User err error + + cachedLinuxUser string + savedLinuxUsers []string + savedLinuxUserErr error } func (m *mockGrantSSHStore) GetCurrentUser() (*entity.User, error) { @@ -102,6 +108,78 @@ func (m *mockGrantSSHStore) ListOrganizations() ([]entity.Organization, error) { return []entity.Organization{*m.org}, nil } +func (m *mockGrantSSHStore) GetCachedLinuxUser() (string, error) { + return m.cachedLinuxUser, nil +} + +func (m *mockGrantSSHStore) SaveCachedLinuxUser(linuxUser string) error { + if m.savedLinuxUserErr != nil { + return m.savedLinuxUserErr + } + m.savedLinuxUsers = append(m.savedLinuxUsers, linuxUser) + return nil +} + +func Test_selectCachedLinuxUser_CachedUserConfirmed(t *testing.T) { + s := &mockGrantSSHStore{cachedLinuxUser: "ubuntu"} + var selectedLabel string + prompter := mockSelector{fn: func(label string, _ []string) string { + selectedLabel = label + return "Yes, use ubuntu" + }} + linuxUser, err := selectCachedLinuxUser(terminal.New(), prompter, s, []string{"ubuntu", "deploy"}) + require.NoError(t, err) + assert.Equal(t, "ubuntu", linuxUser) + assert.Contains(t, selectedLabel, "ubuntu", "cached user must be pre-selected") + assert.Empty(t, s.savedLinuxUsers, "no re-save when the cached user is confirmed") +} + +func Test_selectCachedLinuxUser_CachedUserRejectedPromptsAndSaves(t *testing.T) { + s := &mockGrantSSHStore{cachedLinuxUser: "ubuntu"} + prompter := mockSelector{fn: func(_ string, items []string) string { + if len(items) > 0 && items[0] == "No, choose another" { + return "No, choose another" + } + return items[len(items)-1] // picker returns last option + }} + linuxUser, err := selectCachedLinuxUser(terminal.New(), prompter, s, []string{"ubuntu", "deploy"}) + require.NoError(t, err) + assert.Equal(t, "deploy", linuxUser) + require.Equal(t, []string{"deploy"}, s.savedLinuxUsers) +} + +func Test_selectCachedLinuxUser_NoCachedUserPromptsAndSaves(t *testing.T) { + s := &mockGrantSSHStore{} + prompter := mockSelector{fn: func(_ string, items []string) string { + return items[0] + }} + linuxUser, err := selectCachedLinuxUser(terminal.New(), prompter, s, []string{"ubuntu", "deploy"}) + require.NoError(t, err) + assert.Equal(t, "ubuntu", linuxUser) + require.Equal(t, []string{"ubuntu"}, s.savedLinuxUsers) +} + +func Test_selectCachedLinuxUser_StaleCacheFallsBackToPicker(t *testing.T) { + s := &mockGrantSSHStore{cachedLinuxUser: "gone-user"} + prompter := mockSelector{fn: func(_ string, items []string) string { + return items[0] + }} + linuxUser, err := selectCachedLinuxUser(terminal.New(), prompter, s, []string{"ubuntu", "deploy"}) + require.NoError(t, err) + assert.Equal(t, "ubuntu", linuxUser) + require.Equal(t, []string{"ubuntu"}, s.savedLinuxUsers) +} + +func Test_selectCachedLinuxUser_SaveErrorIsNonFatal(t *testing.T) { + s := &mockGrantSSHStore{savedLinuxUserErr: fmt.Errorf("disk full")} + prompter := mockSelector{fn: func(_ string, items []string) string { + return items[0] + }} + linuxUser, err := selectCachedLinuxUser(terminal.New(), prompter, s, []string{"ubuntu"}) + require.NoError(t, err, "cache save failure must not fail the grant") + assert.Equal(t, "ubuntu", linuxUser) +} + func (m *mockGrantSSHStore) GetOrganizationsByName(name string) ([]entity.Organization, error) { if m.org != nil && m.org.Name == name { return []entity.Organization{*m.org}, nil