diff --git a/pkg/auth/auth_test.go b/pkg/auth/auth_test.go index 4c271b0b9..a1945286b 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 return nil } diff --git a/pkg/cmd/deregister/deregister.go b/pkg/cmd/deregister/deregister.go index efd9090b6..ed1f0f74c 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,65 @@ 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 +} + +// findNodeByDeviceID scans all pages of the org's node list for a node with +// the given device ID. Page size 100 matches ListOrganizationMembers. +func findNodeByDeviceID(ctx context.Context, s externalnode.TokenProvider, deps deregisterDeps, orgID, deviceID string) (string, error) { + client := deps.nodeClients.NewNodeClient(s, config.GlobalConfig.GetBrevPublicAPIURL()) + pageToken := "" + for { + resp, err := client.ListNodes(ctx, connect.NewRequest(&nodev1.ListNodesRequest{ + OrganizationId: orgID, + PageParams: &nodev1.PageParams{ + PageSize: 100, + PageToken: pageToken, + }, + })) + 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 + } + } + pageToken = resp.Msg.GetNextPageToken() + if pageToken == "" { + 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 +163,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 +215,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 95c5ac999..a9ec0241d 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,195 @@ 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(t *testing.T) { + const deviceID = "dev-uuid-pending" + reg := ®ister.DeviceRegistration{ + DisplayName: "My Spark", + OrgID: "org_123", + DeviceID: deviceID, + Status: register.RegistrationStatusPending, + } + + tests := []struct { + name string + listNodesFn func(*nodev1.ListNodesRequest) (*nodev1.ListNodesResponse, error) + removeNodeFn func(*nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) + wantRemovedID string // empty = RemoveNode must not be called + }{ + { + name: "no backend node matches: cleans up locally without RemoveNode", + listNodesFn: func(*nodev1.ListNodesRequest) (*nodev1.ListNodesResponse, error) { + return &nodev1.ListNodesResponse{}, nil + }, + removeNodeFn: func(req *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { + return nil, fmt.Errorf("RemoveNode should not be called with empty ID %q", req.GetExternalNodeId()) + }, + }, + { + name: "backend node recovered by device ID: removed", + listNodesFn: func(req *nodev1.ListNodesRequest) (*nodev1.ListNodesResponse, error) { + return &nodev1.ListNodesResponse{ + Items: []*nodev1.ExternalNode{ + {ExternalNodeId: "unode_other", DeviceId: "dev-different"}, + {ExternalNodeId: "unode_recovered", DeviceId: deviceID}, + }, + }, nil + }, + removeNodeFn: func(req *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { + return &nodev1.RemoveNodeResponse{}, nil + }, + wantRemovedID: "unode_recovered", + }, + { + name: "ListNodes failure is non-fatal: local state still deleted", + listNodesFn: func(*nodev1.ListNodesRequest) (*nodev1.ListNodesResponse, error) { + return nil, connect.NewError(connect.CodeInternal, nil) + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + regStore := &mockRegistrationStore{reg: reg} + 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 tt.listNodesFn(req) + }, + removeNodeFn: func(req *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { + removedNodeID = req.GetExternalNodeId() + return tt.removeNodeFn(req) + }, + } + + 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 tt.wantRemovedID == "" { + if removedNodeID != "" { + t.Errorf("RemoveNode should not be called, got %q", removedNodeID) + } + } else if removedNodeID != tt.wantRemovedID { + t.Errorf("expected RemoveNode called with %q, got %q", tt.wantRemovedID, removedNodeID) + } + + exists, _ := regStore.Exists() + if exists { + t.Error("expected local registration to be deleted") + } + }) + } +} + +// findNodeByDeviceID must walk every page: the pending node can live beyond +// the first ListNodes page. +func Test_runDeregister_PendingRegistration_FoundOnLaterPage(t *testing.T) { + const deviceID = "dev-uuid-pending" + reg := ®ister.DeviceRegistration{ + DisplayName: "My Spark", + OrgID: "org_123", + DeviceID: deviceID, + Status: register.RegistrationStatusPending, + } + regStore := &mockRegistrationStore{reg: reg} + store := &mockDeregisterStore{user: &entity.User{ID: "user_1"}, token: "tok"} + + var removedNodeID string + var pageSize int32 + var pageCount int + svc := &fakeNodeService{ + listNodesFn: func(req *nodev1.ListNodesRequest) (*nodev1.ListNodesResponse, error) { + pageCount++ + pageSize = req.GetPageParams().GetPageSize() + switch req.GetPageParams().GetPageToken() { + case "": + return &nodev1.ListNodesResponse{ + Items: []*nodev1.ExternalNode{{ExternalNodeId: "unode_page1", DeviceId: "dev-other"}}, + NextPageToken: "page-2", + }, nil + case "page-2": + return &nodev1.ListNodesResponse{ + Items: []*nodev1.ExternalNode{{ExternalNodeId: "unode_page2", DeviceId: deviceID}}, + NextPageToken: "", + }, nil + default: + return nil, fmt.Errorf("unexpected page token %q", req.GetPageParams().GetPageToken()) + } + }, + 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 pageCount != 2 { + t.Errorf("expected 2 ListNodes calls, got %d", pageCount) + } + if pageSize != 100 { + t.Errorf("expected page size 100, got %d", pageSize) + } + if removedNodeID != "unode_page2" { + t.Errorf("expected node from page 2 removed, got %q", removedNodeID) + } +} + 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 9788b0e6f..42d0b0388 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) } @@ -96,7 +96,6 @@ func enableSSH( linuxUsername := linuxUser.Username checkSSHDaemon(t) - t.Vprint("") t.Vprint(t.Green("Enabling SSH access on this device")) t.Vprint("") diff --git a/pkg/cmd/grantssh/grantssh.go b/pkg/cmd/grantssh/grantssh.go index 35985834e..8d49a274a 100644 --- a/pkg/cmd/grantssh/grantssh.go +++ b/pkg/cmd/grantssh/grantssh.go @@ -183,12 +183,11 @@ 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)") } + t.Vprint("") + linuxUser = deps.prompter.Select("Select Linux user on the node", linuxUserOptions) } else { selectedUser, err = findUserByIDOrEmail(orgMembers, opts.userIDOrEmail) if err != nil { diff --git a/pkg/cmd/grantssh/grantssh_test.go b/pkg/cmd/grantssh/grantssh_test.go index 10b3cbab0..e1e6441a7 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/ls/ls_test.go b/pkg/cmd/ls/ls_test.go index 0574e4e06..94ca1f9d9 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 315dfb99f..9ec50b3d5 100644 --- a/pkg/cmd/register/device_registration_store.go +++ b/pkg/cmd/register/device_registration_store.go @@ -20,6 +20,11 @@ const ( globalRegistrationDir = "/etc/brev" ) +const ( + RegistrationStatusPending = "pending" + RegistrationStatusRegistered = "registered" +) + // DeviceRegistration is the persistent identity file for a registered device. // Fields align with the AddNodeResponse from dev-plane. type DeviceRegistration struct { @@ -30,17 +35,17 @@ type DeviceRegistration struct { DeviceID string `json:"device_id"` RegisteredAt string `json:"registered_at"` HardwareProfile HardwareProfile `json:"hardware_profile"` + 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 @@ -72,8 +77,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 +90,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 39d7b1a21..2e1b085f0 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 2ad88b435..5a4ee5578 100644 --- a/pkg/cmd/register/register.go +++ b/pkg/cmd/register/register.go @@ -5,7 +5,6 @@ import ( "context" "errors" "fmt" - "os/user" "strings" "time" @@ -78,27 +77,32 @@ func defaultRegisterDeps() registerDeps { } } +type OrgLister interface { + ListOrganizations() ([]entity.Organization, error) +} + 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. 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; both are required. + Use for scripts/CI. +` registerExample = ` # Interactive (prompts for device name, org, confirmations) brev register - # Non-interactive (any flag implies no prompts; --name and --org required) - brev register --name my-node --org my-org - brev register --name my-node --org my-org --ssh-port 22` + # Non-interactive (--name and --org required) + brev register --name my-node --org my-org` ) func NewCmdRegister(t *terminal.Terminal, store RegisterStore) *cobra.Command { var orgFlag string var nameFlag string - var sshPort int + var sshPort int // deprecated var approveFlag bool cmd := &cobra.Command{ @@ -115,7 +119,6 @@ func NewCmdRegister(t *terminal.Terminal, store RegisterStore) *cobra.Command { interactive: interactive, name: nameFlag, orgName: orgFlag, - sshPort: int32(sshPort), skipConfirm: approveFlag, } return runRegister(cmd.Context(), t, store, opts, defaultRegisterDeps()) @@ -126,22 +129,20 @@ 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().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 } -// 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 +150,45 @@ 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) } + if !opts.interactive { if opts.name == "" || opts.orgName == "" { return fmt.Errorf("in non-interactive mode --name and --org are required") } } - - // 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 + if !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 { - return checkExistingRegistration(ctx, t, s, deps) + 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, reg) } - // Capture the device name var name string if opts.interactive { t.Vprint("") @@ -187,13 +205,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 +244,19 @@ func runRegister(ctx context.Context, t *terminal.Terminal, s RegisterStore, opt } } - // Perform the registration steps - reg, err := runRegisterSteps(ctx, t, s, name, org, deps) - 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))) - } - } - - return nil + // Generate the device ID here so a retry reuses it (AddNode is idempotent on device_id). + deviceID := uuid.New().String() + return runRegisterSteps(ctx, t, s, name, org, deps, deviceID) } -// 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 +func runRegisterSteps(ctx context.Context, t *terminal.Terminal, s RegisterStore, name string, org *entity.Organization, deps registerDeps, deviceID 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 +264,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 +273,23 @@ 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, + HardwareProfile: *hwProfile, + Status: RegistrationStatusPending, + RegisteredAt: time.Now().UTC().Format(time.RFC3339), + } + 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,13 +298,13 @@ 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() @@ -311,12 +316,13 @@ func runRegisterSteps(ctx context.Context, t *terminal.Terminal, s RegisterStore DeviceID: deviceID, 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 +331,10 @@ 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 + + t.Vprint("") + t.Vprintf(" %s\n", t.Green("To enable SSH access to this device, run: brev enable-ssh")) + return nil } func resolveOrgInteractive(t *terminal.Terminal, s RegisterStore, deps registerDeps) (*entity.Organization, error) { @@ -348,26 +357,25 @@ 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 checkExistingRegistration(ctx context.Context, t *terminal.Terminal, s RegisterStore, deps registerDeps) error { - reg, loadErr := deps.registrationStore.Load() - if loadErr != nil { - return fmt.Errorf("this machine is already registered but the registration file could not be read: %w", loadErr) +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, reg *DeviceRegistration) error { t.Vprint("") t.Vprintf(" This machine is already registered as %s (%s).\n", reg.DisplayName, reg.ExternalNodeID) 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 +432,21 @@ 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) - if err != nil { - return fmt.Errorf("allocate SSH port failed: %w", err) - } - - err = SetupAndRegisterNodeSSHAccess(ctx, t, deps.nodeClients, tokenProvider, reg, brevUser, osUser.Username, brevPortID) - if err != nil { - return fmt.Errorf("grant SSH failed: %w", err) - } - - t.Vprint("") - t.Vprint(t.Green(fmt.Sprintf("SSH access enabled. You can now SSH to this device via: brev shell %s", reg.DisplayName))) - t.Vprint("") - return nil + org := &entity.Organization{ID: pending.OrgID, Name: pending.OrgName} + return runRegisterSteps(ctx, t, s, pending.DisplayName, org, deps, pending.DeviceID) } diff --git a/pkg/cmd/register/register_test.go b/pkg/cmd/register/register_test.go index d98b1a92a..8b4f4161c 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" @@ -72,7 +73,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 +185,28 @@ func testRegisterDeps(t *testing.T, svc *fakeNodeService, regStore RegistrationS }, server } +func testRegisterStore() *mockRegisterStore { + return &mockRegisterStore{ + user: &entity.User{ID: "user_1"}, + org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, + token: "tok", + } +} + +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 +237,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 +253,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 +283,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 +302,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 +376,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 +385,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,153 +399,148 @@ 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{} - - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_default", Name: "DefaultOrg"}, - orgs: []entity.Organization{ - {ID: "org_456", Name: "SpecificOrg"}, + tests := []struct { + name string + orgs []entity.Organization + orgName string + wantErr string + wantOrgID string + }{ + { + name: "resolves org by name", + orgs: []entity.Organization{{ID: "org_456", Name: "SpecificOrg"}}, + orgName: "SpecificOrg", + wantOrgID: "org_456", }, - token: "tok", - } - - var capturedOrgID string - svc := &fakeNodeService{ - addNodeFn: func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { - capturedOrgID = req.GetOrganizationId() - return &nodev1.AddNodeResponse{ - ExternalNode: &nodev1.ExternalNode{ - ExternalNodeId: "unode_abc", - OrganizationId: req.GetOrganizationId(), - Name: req.GetName(), - DeviceId: req.GetDeviceId(), - }, - }, nil + { + name: "org not found", + orgs: []entity.Organization{}, + orgName: "NonexistentOrg", + wantErr: "no organization found", }, } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + regStore := &mockRegistrationStore{} - setupRunner := &mockSetupRunner{} - deps, server := testRegisterDeps(t, svc, regStore) - defer server.Close() - deps.setupRunner = setupRunner - - SetTestSSHPort(22) - defer ClearTestSSHPort() - - term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "SpecificOrg", sshPort: 22} - err := runRegister(context.Background(), term, store, opts, deps) - if err != nil { - t.Fatalf("runRegister with --org failed: %v", err) - } - - if capturedOrgID != "org_456" { - t.Errorf("expected org_456, got %s", capturedOrgID) - } + store := &mockRegisterStore{ + user: &entity.User{ID: "user_1"}, + org: &entity.Organization{ID: "org_default", Name: "DefaultOrg"}, + orgs: tt.orgs, + token: "tok", + } - reg, err := regStore.Load() - if err != nil { - t.Fatalf("Load failed: %v", err) - } - if reg.OrgID != "org_456" { - t.Errorf("expected registration org org_456, got %s", reg.OrgID) - } -} + var capturedOrgID string + svc := &fakeNodeService{ + addNodeFn: func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { + capturedOrgID = req.GetOrganizationId() + return &nodev1.AddNodeResponse{ + ExternalNode: &nodev1.ExternalNode{ + ExternalNodeId: "unode_abc", + OrganizationId: req.GetOrganizationId(), + Name: req.GetName(), + DeviceId: req.GetDeviceId(), + }, + }, nil + }, + } -func Test_runRegister_WithOrgFlag_NotFound(t *testing.T) { - regStore := &mockRegistrationStore{} + setupRunner := &mockSetupRunner{} + deps, server := testRegisterDeps(t, svc, regStore) + defer server.Close() + deps.setupRunner = setupRunner - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_default", Name: "DefaultOrg"}, - orgs: []entity.Organization{}, - token: "tok", - } + term := terminal.New() + opts := registerOpts{interactive: false, name: "my-spark", orgName: tt.orgName} + err := runRegister(context.Background(), term, store, opts, deps) + if tt.wantErr != "" { + if err == nil { + t.Fatal("expected error when org not found") + } + if !strings.Contains(err.Error(), tt.wantErr) { + t.Errorf("expected %q error, got: %v", tt.wantErr, err) + } + return + } + if err != nil { + t.Fatalf("runRegister with --org failed: %v", err) + } - svc := &fakeNodeService{} - deps, server := testRegisterDeps(t, svc, regStore) - defer server.Close() + if capturedOrgID != tt.wantOrgID { + t.Errorf("expected org %s, got %s", tt.wantOrgID, capturedOrgID) + } - term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "NonexistentOrg", sshPort: 22} - err := runRegister(context.Background(), term, store, opts, deps) - if err == nil { - t.Fatal("expected error when org not found") - } - if !strings.Contains(err.Error(), "no organization found") { - t.Errorf("expected 'no organization found' error, got: %v", err) + reg, err := regStore.Load(false) + if err != nil { + t.Fatalf("Load failed: %v", err) + } + if reg.OrgID != tt.wantOrgID { + t.Errorf("expected registration org %s, got %s", tt.wantOrgID, reg.OrgID) + } + }) } } -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 +563,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,147 +656,97 @@ Peers count: 0/0 Connected` } } -func Test_runRegister_GrantSSH_retries_on_connection_error_then_succeeds(t *testing.T) { - regStore := &mockRegistrationStore{} - - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - token: "tok", +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) - var grantCalls 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 - }, - grantNodeSSHAccessFn: func(_ *nodev1.GrantNodeSSHAccessRequest) (*nodev1.GrantNodeSSHAccessResponse, error) { - grantCalls++ - if grantCalls < 2 { - return nil, connect.NewError(connect.CodeInternal, nil) + 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") } - 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) - 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 !strings.Contains(err.Error(), tt.errSubstr) { + t.Errorf("expected error containing %q, got: %v", tt.errSubstr, err) + } + }) } } -func Test_runRegister_GrantSSH_no_retry_on_permanent_error(t *testing.T) { +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", - } - - var grantCalls 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 - }, - grantNodeSSHAccessFn: func(_ *nodev1.GrantNodeSSHAccessRequest) (*nodev1.GrantNodeSSHAccessResponse, error) { - grantCalls++ - return nil, connect.NewError(connect.CodePermissionDenied, nil) - }, - } + store := testRegisterStore() + svc := &fakeNodeService{} 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: "", orgName: ""} 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) + if err == nil { + t.Fatal("expected error when no name/org in non-interactive mode") } - - if grantCalls != 1 { - t.Errorf("expected GrantNodeSSHAccess to be called once (no retry on permanent error), got %d", grantCalls) + 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_NameValidation(t *testing.T) { +func Test_runRegister_ResumesPendingRegistration(t *testing.T) { + const pendingDeviceID = "device-uuid-pending" tests := []struct { - name string - input string - wantErr bool - errSubstr string + name string + opts registerOpts + store func() *mockRegisterStore }{ - {"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 + { + name: "interactive mode", + opts: registerOpts{interactive: true}, + store: testRegisterStore, + }, + { + name: "non-interactive with matching --org", + opts: registerOpts{interactive: false, name: "My Spark", orgName: "TestOrg"}, + store: func() *mockRegisterStore { + return &mockRegisterStore{ + user: &entity.User{ID: "user_1"}, + orgs: []entity.Organization{{ID: "org_123", Name: "TestOrg"}}, + token: "tok", + } + }, + }, } - 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", - } + regStore := &mockRegistrationStore{reg: testPendingReg("org_123", "TestOrg", pendingDeviceID)} + store := tt.store() + var addNodeDeviceIDs []string svc := &fakeNodeService{ addNodeFn: func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { + 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{ + RegistrationCommand: "netbird up --key abc", + }, }, }, nil }, @@ -811,321 +755,130 @@ func Test_runRegister_NameValidation(t *testing.T) { 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 err := runRegister(context.Background(), term, store, tt.opts, deps); err != nil { + t.Fatalf("runRegister failed: %v", err) } - }) - } -} - -func Test_runRegister_PlatformIncompatible(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{} - deps, server := testRegisterDeps(t, svc, regStore) - defer server.Close() - - deps.platform = mockPlatform{compatible: false} + if len(addNodeDeviceIDs) != 1 || addNodeDeviceIDs[0] != pendingDeviceID { + t.Errorf("expected AddNode to reuse device ID %q once, got %v", pendingDeviceID, addNodeDeviceIDs) + } - 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) + 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) + } + }) } } -func Test_runRegister_HardwareProfilerFailure(t *testing.T) { - regStore := &mockRegistrationStore{} +// --- Org mismatch --- - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - token: "tok", +func Test_runRegister_OrgMismatch(t *testing.T) { + tests := []struct { + name string + status string + useOrgFlag bool // pass --org for the new org + wantWording string // pending -> "incomplete registration"; registered -> "already registered" + }{ + {"OrgFlag_Pending", RegistrationStatusPending, true, "incomplete registration"}, + {"OrgFlag_AlreadyRegistered", RegistrationStatusRegistered, true, "already registered"}, } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + regStore := &mockRegistrationStore{reg: &DeviceRegistration{ + ExternalNodeID: "unode_existing", + DisplayName: "My Spark", + OrgID: "org_other", + OrgName: "OtherOrg", + DeviceID: "dev-pending", + Status: tt.status, + }} + store := &mockRegisterStore{ + user: &entity.User{ID: "user_1"}, + orgs: []entity.Organization{{ID: "org_123", Name: "TestOrg"}}, // new org + token: "tok", + } - svc := &fakeNodeService{} - deps, server := testRegisterDeps(t, svc, regStore) - defer server.Close() + var addNodeCalls int + svc := &fakeNodeService{ + addNodeFn: func(_ *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { + addNodeCalls++ + return nil, fmt.Errorf("AddNode should not be called on org mismatch") + }, + } - deps.hardwareProfiler = &mockHardwareProfiler{err: fmt.Errorf("nvml init failed")} + 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 hardware profiler fails") - } - if !strings.Contains(err.Error(), "hardware profile") { - t.Errorf("expected hardware profile error, got: %v", err) + 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 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_NetBirdInstallFailure(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{} - deps, server := testRegisterDeps(t, svc, regStore) - defer server.Close() - - deps.netbird = mockNetBirdManager{err: fmt.Errorf("install failed")} +func Test_runRegister_ResumeAddNodeFails_StaysPending(t *testing.T) { + const pendingDeviceID = "device-uuid-pending" + regStore := &mockRegistrationStore{reg: testPendingReg("org_123", "TestOrg", pendingDeviceID)} - 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) - } -} + store := testRegisterStore() -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", + var addNodeCalls int + svc := &fakeNodeService{ + addNodeFn: func(_ *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { + addNodeCalls++ + return nil, connect.NewError(connect.CodeInternal, nil) + }, } - svc := &fakeNodeService{} deps, server := testRegisterDeps(t, svc, regStore) defer server.Close() term := terminal.New() - opts := registerOpts{interactive: false, name: "", orgName: "", sshPort: 22} - err := runRegister(context.Background(), term, store, opts, deps) + err := runRegister(context.Background(), term, store, registerOpts{interactive: true}, 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) + t.Fatal("expected error when AddNode fails during resume") } -} - -func Test_runRegister_NoNameAlreadyRegistered(t *testing.T) { - regStore := &mockRegistrationStore{ - reg: &DeviceRegistration{ - ExternalNodeID: "unode_existing", - DisplayName: "Existing Device", - OrgID: "org_123", - }, - } - - store := &mockRegisterStore{ - user: &entity.User{ID: "user_1"}, - org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, - token: "tok", + if addNodeCalls != 1 { + t.Fatalf("expected AddNode called once, got %d", addNodeCalls) } - svc := &fakeNodeService{ - getNodeFn: func(req *nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) { - return &nodev1.GetNodeResponse{ - ExternalNode: &nodev1.ExternalNode{ - ExternalNodeId: req.GetExternalNodeId(), - ConnectivityInfo: &nodev1.ConnectivityInfo{ - Status: nodev1.NetworkMemberStatus_NETWORK_MEMBER_STATUS_CONNECTED, - }, - }, - }, nil - }, + reg, loadErr := regStore.Load(true) + if loadErr != nil { + t.Fatalf("Load failed: %v", loadErr) } - - deps, server := testRegisterDeps(t, svc, regStore) - defer server.Close() - - term := terminal.New() - opts := registerOpts{interactive: false, name: "Existing", orgName: "TestOrg", sshPort: 22} - 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 reg.Status != RegistrationStatusPending { + t.Errorf("expected record to stay pending after failed resume, got %q", reg.Status) } - - // Registration should still exist - exists, _ := regStore.Exists() - if !exists { - t.Error("expected registration to still exist") + if reg.DeviceID != pendingDeviceID { + t.Errorf("expected device ID to remain %q, got %q", pendingDeviceID, reg.DeviceID) } -} - -func Test_runRegister_OpenSSHPort(t *testing.T) { // nolint:funlen, gocyclo, gocognit // test - 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: "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()) - } - }, - }, - } - - 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", - } - - var gotOpenReq *nodev1.OpenPortRequest - var gotGrantReq *nodev1.GrantNodeSSHAccessRequest - 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 - }, - } - - 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} - err := runRegister(context.Background(), term, store, opts, deps) - - tt.verify(t, gotOpenReq, gotGrantReq, regStore, err) - }) + if reg.ExternalNodeID != "" { + t.Errorf("expected no ExternalNodeID after failed resume, got %q", reg.ExternalNodeID) } } diff --git a/pkg/cmd/register/sshkeys.go b/pkg/cmd/register/sshkeys.go index 1188766dc..2ca9bb910 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 447c53cb4..71165928c 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") }