diff --git a/.agents/skills/brev-cli/SKILL.md b/.agents/skills/brev-cli/SKILL.md index f054e4526..5766acb69 100644 --- a/.agents/skills/brev-cli/SKILL.md +++ b/.agents/skills/brev-cli/SKILL.md @@ -202,6 +202,34 @@ brev ls --json | jq -r '.workspaces[].name' brev ls nodes --json | jq -r '.[] | select(.status=="Connected") | .name' ``` +### BYON Network and SSH + +For a machine you bring to Brev, network membership and SSH credentials are +separate operations: + +```bash +# Join only the organization's Brev/NetBird network. +brev join + +# Optionally enable access for yourself, then grant a collaborator. +brev enable-ssh +brev grant-ssh + +# Explicitly revoke tracked SSH grants, then retire membership. +brev disable-ssh +brev leave +``` + +`brev register` and `brev deregister` are deprecated aliases for `join` and +`leave`; they warn when executed. `enable-ssh` requires an existing join and +can reconnect its tunnel, but never joins a network. Use `grant-ssh` and +`revoke-ssh` for individual collaborators. `disable-ssh` best-effort revokes +every backend-tracked SSH grant across the node, continuing after individual +failures and revoking the invoking Brev user's own access last. It does not +modify `authorized_keys`, close ports, stop `sshd`, end active sessions, or +leave the network. `leave` removes membership without running that per-grant +revocation flow or cleaning up local keys. + ### Instance Management ```bash # List instances diff --git a/.agents/skills/brev-cli/reference/commands.md b/.agents/skills/brev-cli/reference/commands.md index fe91cbbd5..86de29d07 100644 --- a/.agents/skills/brev-cli/reference/commands.md +++ b/.agents/skills/brev-cli/reference/commands.md @@ -497,6 +497,93 @@ Generate an invite link. brev invite ``` +## BYON Network and SSH Commands + +These commands apply to a machine brought into a Brev organization. Network +membership and Brev-managed SSH credentials are separate. + +### Canonical workflows + +```bash +# Join networking only. Then optionally enable your SSH access and grant a collaborator. +brev join +brev enable-ssh +brev grant-ssh + +# Explicitly revoke tracked SSH grants before retiring network membership. +brev disable-ssh +brev leave +``` + +### brev join / brev register + +Join a device to the organization's Brev/NetBird network. + +```bash +brev join [--name --org ] [--approve] +``` + +`join` establishes membership only: it does not enable SSH or allocate an SSH +port. `register` is a deprecated alias that warns on execution. The old +`--ssh-port` flag is no longer supported; migrate scripts to `brev join` and +then `brev enable-ssh` on the joined machine. + +### brev enable-ssh + +Enable Brev-managed SSH for the invoking Brev user on the joined node. + +```bash +brev enable-ssh +``` + +This requires an existing `join`. It confirms the existing Brev tunnel and can +reconnect it when disconnected, but it does not add a node, select an +organization, save a registration, or join a network. + +### brev grant-ssh / brev revoke-ssh + +Manage an individual collaborator's SSH access tuple on a node. + +```bash +brev grant-ssh +brev revoke-ssh +``` + +Use these commands for collaborator access rather than treating `enable-ssh` +or `disable-ssh` as collaborator-management commands. + +### brev disable-ssh + +Revoke all backend-tracked Brev SSH grants from the joined node. + +```bash +brev disable-ssh [--approve] +``` + +This node-wide operation makes a best-effort attempt to revoke each exact active +backend access tuple. It continues after individual failures, reports an error +when any tuple remains, and revokes the invoking Brev user's own access last. It +does not inspect or modify local `authorized_keys` files. It leaves existing +ports allocated, leaves `sshd` running, does not forcibly terminate active SSH +sessions, and does not remove membership or the backend node. + +### brev leave / brev deregister + +Remove Brev network membership from the device. + +```bash +brev leave [--approve] +``` + +`leave` removes the backend node, VPN route, and local registration. It does not +run `disable-ssh`'s per-grant revocation flow or modify local `authorized_keys` +files. Run `brev disable-ssh` first when explicit best-effort revocation of +tracked grants is desired. `deregister` is a deprecated alias that warns on +execution. + +`leave` continues to uninstall NetBird even if it was installed before Brev. +Install-ownership tracking is a follow-up, so ensure that removal is intended. + ## Configuration Commands ### brev login / brev logout diff --git a/.gitignore b/.gitignore index 1276cd3d7..2d8331945 100644 --- a/.gitignore +++ b/.gitignore @@ -55,4 +55,8 @@ devworkspace/** test.txt test2.txt homebrew-brev -flake-explorations \ No newline at end of file +flake-explorations + +# AI +/docs/superpowers/plans/ +/docs/superpowers/specs/ \ No newline at end of file diff --git a/CHANGELOG.md b/CHANGELOG.md index 486480540..55bf7a8ba 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -10,3 +10,18 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Added [WIP] Add tailscale vpn client embedded with Brev. + +- `brev join` for Brev/NetBird network membership, `brev leave` for membership teardown, and node-wide `brev disable-ssh`. + +### Changed + +- `brev join` no longer enables SSH. `brev enable-ssh` requires an existing joined membership and reconnects its tunnel when needed. +- `brev disable-ssh` best-effort revokes all backend-tracked SSH grants, revokes the invoking user's own access last, and no longer sweeps local `authorized_keys` files. + +### Deprecated + +- `brev register` and `brev deregister` remain compatibility aliases for `join` and `leave`, and warn on stderr when executed. + +### Migration + +- Scripts using `--ssh-port` must run `brev join` followed by `brev enable-ssh`. diff --git a/README.md b/README.md index f625c153a..7eda82f7c 100644 --- a/README.md +++ b/README.md @@ -76,6 +76,8 @@ brev ls https://docs.nvidia.com/brev/latest/ +[Bring Your Own Node (BYON) network and SSH workflows](docs/BYON.md) + --- ## AI Agent Integration diff --git a/docs/BYON.md b/docs/BYON.md new file mode 100644 index 000000000..c9807a825 --- /dev/null +++ b/docs/BYON.md @@ -0,0 +1,71 @@ +# Bring Your Own Node (BYON) + +Brev separates network membership from Brev-managed SSH credentials on a +machine you bring to your organization. + +## Join networking only + +```bash +brev join +``` + +`brev join` establishes this machine's Brev/NetBird organization membership. +It does not enable SSH or create an SSH port. Use `brev register` only for +compatibility with existing automation: it is a deprecated alias for `join` and +prints a warning when executed. + +Scripts that used `--ssh-port` must migrate to two commands: + +```bash +brev join +brev enable-ssh +``` + +## Enable and grant SSH + +After joining, enable Brev-managed SSH for the invoking Brev user: + +```bash +brev enable-ssh +``` + +`enable-ssh` requires a prior join. It confirms the existing Brev tunnel and +can reconnect it when it is disconnected; it never joins a network or creates +membership. It then enables the invoking user's access on the joined node. + +Grant and revoke collaborator access separately: + +```bash +brev grant-ssh +brev revoke-ssh +``` + +These commands manage individual collaborator access tuples. They are not part +of `join`, `enable-ssh`, or the node-wide revocation command. + +## Retire Brev access and membership + +To explicitly revoke Brev-tracked SSH grants before leaving the network: + +```bash +brev disable-ssh +brev leave +``` + +`disable-ssh` is node-wide. It makes a best-effort attempt to revoke every +backend-tracked Brev SSH access tuple, continuing after individual failures and +revoking the invoking Brev user's own access last. It returns an error if any +revocation fails so the remaining records can be retried. It does not inspect or +modify local `authorized_keys` files. It leaves existing ports allocated, leaves +`sshd` running, does not forcibly terminate active SSH sessions, and does not +change network membership. + +`leave` removes the backend node, Brev VPN route, and local registration. It +does not run the per-grant `disable-ssh` flow or modify local `authorized_keys` +files. Run `disable-ssh` first when explicit best-effort revocation of tracked +grants is desired. `brev deregister` is a deprecated alias for `leave` and warns +when executed. + +`leave` preserves the existing behavior of uninstalling NetBird even when +NetBird was installed before Brev. Tracking whether Brev owns that installation +is a follow-up improvement, so use `leave` only when that removal is intended. diff --git a/pkg/cmd/cmd.go b/pkg/cmd/cmd.go index 49b5254d7..439faab5f 100644 --- a/pkg/cmd/cmd.go +++ b/pkg/cmd/cmd.go @@ -15,6 +15,7 @@ import ( "github.com/brevdev/brev-cli/pkg/cmd/copy" "github.com/brevdev/brev-cli/pkg/cmd/delete" "github.com/brevdev/brev-cli/pkg/cmd/deregister" + "github.com/brevdev/brev-cli/pkg/cmd/disablessh" "github.com/brevdev/brev-cli/pkg/cmd/enablessh" "github.com/brevdev/brev-cli/pkg/cmd/envvars" "github.com/brevdev/brev-cli/pkg/cmd/exec" @@ -316,10 +317,11 @@ func createCmdTree(cmd *cobra.Command, t *terminal.Terminal, loginCmdStore *stor cmd.AddCommand(reset.NewCmdReset(t, loginCmdStore, noLoginCmdStore)) cmd.AddCommand(profile.NewCmdProfile(t, loginCmdStore, noLoginCmdStore)) cmd.AddCommand(refresh.NewCmdRefresh(t, loginCmdStore)) - cmd.AddCommand(register.NewCmdRegister(t, externalNodeCmdStore)) - cmd.AddCommand(deregister.NewCmdDeregister(t, externalNodeCmdStore)) + cmd.AddCommand(register.NewCmdJoin(t, externalNodeCmdStore)) + cmd.AddCommand(deregister.NewCmdLeave(t, externalNodeCmdStore)) cmd.AddCommand(upgrade.NewCmdUpgrade(t, noLoginCmdStore)) cmd.AddCommand(enablessh.NewCmdEnableSSH(t, externalNodeCmdStore)) + cmd.AddCommand(disablessh.NewCmdDisableSSH(t, externalNodeCmdStore)) cmd.AddCommand(grantssh.NewCmdGrantSSH(t, externalNodeCmdStore)) cmd.AddCommand(revokessh.NewCmdRevokeSSH(t, externalNodeCmdStore)) cmd.AddCommand(runtasks.NewCmdRunTasks(t, noLoginCmdStore)) diff --git a/pkg/cmd/deregister/deregister.go b/pkg/cmd/deregister/deregister.go index efd9090b6..b12af73d5 100644 --- a/pkg/cmd/deregister/deregister.go +++ b/pkg/cmd/deregister/deregister.go @@ -1,18 +1,20 @@ -// Package deregister provides the brev deregister command for device deregistration +// Package deregister provides the canonical Brev network leave command and +// its deprecated deregister alias. package deregister import ( "context" "fmt" - "os/user" + "io" + nodev1connect "buf.build/gen/go/brevdev/devplane/connectrpc/go/devplaneapi/v1/devplaneapiv1connect" nodev1 "buf.build/gen/go/brevdev/devplane/protocolbuffers/go/devplaneapi/v1" "connectrpc.com/connect" - breverrors "github.com/brevdev/brev-cli/pkg/errors" "github.com/brevdev/brev-cli/pkg/cmd/register" "github.com/brevdev/brev-cli/pkg/config" "github.com/brevdev/brev-cli/pkg/entity" + breverrors "github.com/brevdev/brev-cli/pkg/errors" "github.com/brevdev/brev-cli/pkg/externalnode" "github.com/brevdev/brev-cli/pkg/sudo" "github.com/brevdev/brev-cli/pkg/terminal" @@ -20,192 +22,188 @@ import ( "github.com/spf13/cobra" ) -// DeregisterStore defines the store methods needed by the deregister command. -type DeregisterStore interface { +// LeaveStore defines the authenticated store methods needed by leave. +type LeaveStore 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{} +// DeregisterStore is retained for source compatibility. +// +// Deprecated: use LeaveStore. +type DeregisterStore = LeaveStore -func (brevSSHKeyRemover) RemoveBrevKeys(u *user.User) ([]string, error) { - removed, err := register.RemoveBrevAuthorizedKeys(u) - if err != nil { - return nil, fmt.Errorf("removing brev authorized keys: %w", err) - } - return removed, nil +type netBirdUninstaller interface { + Uninstall() error } -// deregisterDeps bundles the side-effecting dependencies of runDeregister so -// they can be replaced in tests. -type deregisterDeps struct { +type leaveDeps struct { platform externalnode.PlatformChecker - prompter terminal.Selector confirmer terminal.Confirmer gater sudo.Gater - netbird register.NetBirdManager + netbird netBirdUninstaller nodeClients externalnode.NodeClientFactory registrationStore register.RegistrationStore - sshKeys SSHKeyRemover } -func defaultDeregisterDeps() deregisterDeps { - return deregisterDeps{ +func defaultLeaveDeps() leaveDeps { + return leaveDeps{ platform: register.LinuxPlatform{}, - prompter: register.TerminalPrompter{}, confirmer: register.TerminalPrompter{}, gater: sudo.Default, netbird: register.Netbird{}, nodeClients: register.DefaultNodeClientFactory{}, registrationStore: register.NewFileRegistrationStore(), - sshKeys: brevSSHKeyRemover{}, } } -var ( - deregisterLong = `Deregister your device from NVIDIA Brev +const leaveLong = `Leave the Brev network -This command removes the local registration data and uninstalls -the Brev tunnel (network agent).` +This removes the backend node, uninstalls the Brev tunnel, and deletes local +registration data. It does not revoke SSH access grants; run "brev disable-ssh" +first when those grants should be revoked.` - deregisterExample = ` brev deregister` -) +// NewCmdLeave creates the canonical network-membership teardown command. +func NewCmdLeave(t *terminal.Terminal, store LeaveStore) *cobra.Command { + return newCmdLeave(t, store, defaultLeaveDeps()) +} +// NewCmdDeregister is retained for source compatibility. It returns the +// canonical leave command with deregister as its deprecated alias. +// +// Deprecated: use NewCmdLeave. func NewCmdDeregister(t *terminal.Terminal, store DeregisterStore) *cobra.Command { - var approveFlag bool + return NewCmdLeave(t, store) +} +func newCmdLeave(t *terminal.Terminal, store LeaveStore, deps leaveDeps) *cobra.Command { + var approveFlag bool cmd := &cobra.Command{ Annotations: map[string]string{"configuration": ""}, - Use: "deregister", + Use: "leave", + Aliases: []string{"deregister"}, DisableFlagsInUseLine: true, - Short: "Deregister your device from Brev", - Long: deregisterLong, - Example: deregisterExample, - RunE: func(cmd *cobra.Command, args []string) error { - return runDeregister(cmd.Context(), t, store, defaultDeregisterDeps(), approveFlag) + Short: "Leave the Brev network", + Long: leaveLong, + Example: " brev leave\n brev leave --approve", + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, _ []string) error { + if cmd.CalledAs() == "deregister" { + _, _ = fmt.Fprintln(cmd.ErrOrStderr(), `Warning: "brev deregister" is deprecated; use "brev leave" instead.`) + _, _ = fmt.Fprintln(cmd.ErrOrStderr(), `This command does not revoke SSH access grants; run "brev disable-ssh" before leaving if you want to revoke them.`) + } + return runLeave(cmd.Context(), t, cmd.ErrOrStderr(), store, deps, approveFlag) }, } - cmd.Flags().BoolVar(&approveFlag, "approve", false, "skip confirmation prompt (assume yes)") - return cmd } -func runDeregister(ctx context.Context, t *terminal.Terminal, s DeregisterStore, deps deregisterDeps, skipConfirm bool) error { //nolint:funlen,gocyclo // deregistration flow +func runLeave( + ctx context.Context, + t *terminal.Terminal, + warnings io.Writer, + store LeaveStore, + deps leaveDeps, + skipConfirm bool, +) error { //nolint:funlen // The retry-safe teardown order is intentionally explicit. if !deps.platform.IsCompatible() { - return fmt.Errorf("brev deregister is only supported on Linux") - } - - if err := deps.gater.Gate(t, deps.confirmer, "Device deregistration", skipConfirm); err != nil { - return fmt.Errorf("sudo issue: %w", err) + return fmt.Errorf("brev leave is only supported on Linux") } reg, err := deps.registrationStore.Load() if err != nil { - return err //nolint:wrapcheck // do not present stack trace for this error + return fmt.Errorf("read joined-device registration: %w", err) } - - // Only prompt for login when there is a device to deregister. - if _, err := s.GetCurrentUser(); err != nil { + if _, err := store.GetCurrentUser(); err != nil { return breverrors.WrapAndTrace(err) } - orgName := reg.OrgName - if orgName == "" { - orgName = "(unknown)" + client := deps.nodeClients.NewNodeClient(store, config.GlobalConfig.GetBrevPublicAPIURL()) + node, missing, err := lookupJoinedNodeForLeave(ctx, client, reg) + if err != nil { + return fmt.Errorf("inspect joined node before leaving: %w", err) } - osUser, _ := user.Current() - linuxUser := "(unknown)" - if osUser != nil { - linuxUser = osUser.Username + if warnings == nil { + warnings = io.Discard + } + _, _ = fmt.Fprintln(warnings, "Leaving removes the Brev tunnel and may interrupt commands using Brev SSH. Run this locally or through out-of-band access.") + if missing { + _, _ = fmt.Fprintln(warnings, "Warning: the backend node is already absent; skipping SSH grant inspection.") + } else { + grantCount, accountCount := remainingSSHAccessCounts(node.GetSshAccess()) + if grantCount > 0 { + _, _ = fmt.Fprintf(warnings, "Warning: %d SSH grants across %d Linux accounts remain on this node.\n", grantCount, accountCount) + _, _ = fmt.Fprintln(warnings, `Leaving stops Brev-routed SSH but does not revoke these grants. Cancel and run "brev disable-ssh" first if you want them revoked.`) + } } t.Vprint("") t.Vprint(t.White("══════════════════════════════════════════════════")) - t.Vprint(t.White(" Deregistering your device from Brev")) + t.Vprint(t.White(" Leaving the Brev network")) t.Vprint(t.White("══════════════════════════════════════════════════")) - t.Vprint("") - if !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(orgName+" ("+reg.OrgID+")")) - t.Vprintf(" %s %s\n", t.Green(fmt.Sprintf("%-14s", "Linux user:")), t.BoldBlue(linuxUser)) - t.Vprint("") - t.Vprint(t.Yellow(" This will:")) - t.Vprint(" 1. Remove this node from Brev") - t.Vprint(" 2. Remove Brev SSH keys from this machine (if any)") - t.Vprint(" 3. Uninstall the Brev tunnel") - t.Vprint(" 4. Delete local registration data") + t.Vprintf(" Node: %s (%s)\n", reg.DisplayName, reg.ExternalNodeID) + t.Vprintf(" Organization: %s (%s)\n", reg.OrgName, reg.OrgID) t.Vprint("") - if !skipConfirm { - confirm := deps.prompter.Select( - "Proceed with deregistration?", - []string{"Yes, proceed", "No, cancel"}, - ) - if confirm != "Yes, proceed" { - t.Vprint("Deregistration canceled.") - return nil - } + if !skipConfirm && !deps.confirmer.ConfirmYesNo("Leave the Brev network?") { + t.Vprint("Leave canceled.") + return nil + } + if err := deps.gater.Gate(t, deps.confirmer, "Leave Brev network", true); err != nil { + return fmt.Errorf("sudo issue: %w", err) } - 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 != nil && connect.CodeOf(err) != connect.CodeNotFound { + return fmt.Errorf("leave Brev network: remove node: %w", err) } - t.Vprintf("%s Node removed from Brev.\n", t.Green(" ✓")) - t.Vprint("") - - t.Vprint(t.Yellow("[Step 2/4] Removing Brev SSH keys...")) - if osUser == nil { - t.Vprintf(" %s\n", t.Yellow("Skipped: could not determine current user")) - } else { - removed, kerr := deps.sshKeys.RemoveBrevKeys(osUser) - switch { - case kerr != nil: - t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: failed to remove Brev SSH keys: %v", kerr))) - case len(removed) > 0: - t.Vprintf("%s Brev SSH keys removed from authorized_keys:\n", t.Green(" ✓")) - for _, key := range removed { - t.Vprintf(" - %s\n", key) - } - default: - t.Vprint(" No Brev SSH keys found in authorized_keys.") - } + if err := deps.netbird.Uninstall(); err != nil { + return fmt.Errorf("leave Brev network: uninstall tunnel: %w", err) } - t.Vprint("") - - t.Vprint(t.Yellow("[Step 3/4] Removing Brev tunnel...")) - err = deps.netbird.Uninstall() - if err != nil { - t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: failed to remove Brev tunnel: %v", err))) - } else { - t.Vprintf("%s Brev tunnel removed.\n", t.Green(" ✓")) + if err := deps.registrationStore.Delete(); err != nil { + return fmt.Errorf("leave Brev network: delete local registration: %w", err) } - t.Vprint("") + t.Vprint("Left the Brev network.") + return nil +} - t.Vprint(t.Yellow("[Step 4/4] Removing registration data...")) - err = deps.registrationStore.Delete() +func lookupJoinedNodeForLeave( + ctx context.Context, + client nodev1connect.ExternalNodeServiceClient, + reg *register.DeviceRegistration, +) (*nodev1.ExternalNode, bool, error) { + resp, err := client.ListNodes(ctx, connect.NewRequest(&nodev1.ListNodesRequest{ + OrganizationId: reg.OrgID, + })) if err != nil { - t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: failed to remove local registration file: %v", err))) - t.Vprint(" You can manually remove it with: rm /etc/brev/device_registration.json") - } else { - t.Vprintf("%s Registration data removed.\n", t.Green(" ✓")) + return nil, false, fmt.Errorf("list organization nodes: %w", err) } - t.Vprintf("%s Deregistration complete.\n", t.Green(" ✓")) - t.Vprint("") + if resp == nil || resp.Msg == nil { + return nil, false, fmt.Errorf("list organization nodes: empty response") + } + for _, candidate := range resp.Msg.GetItems() { + if candidate != nil && candidate.GetExternalNodeId() == reg.ExternalNodeID { + return candidate, false, nil + } + } + if resp.Msg.GetNextPageToken() != "" { + return nil, false, fmt.Errorf("registered node was not in the returned page and node listing is incomplete") + } + return nil, true, nil +} - return nil +func remainingSSHAccessCounts(accesses []*nodev1.SSHAccess) (int, int) { + accounts := make(map[string]struct{}, len(accesses)) + grantCount := 0 + for _, access := range accesses { + if access == nil { + continue + } + grantCount++ + accounts[access.GetLinuxUser()] = struct{}{} + } + return grantCount, len(accounts) } diff --git a/pkg/cmd/deregister/deregister_test.go b/pkg/cmd/deregister/deregister_test.go index 95c5ac999..757d2bb8c 100644 --- a/pkg/cmd/deregister/deregister_test.go +++ b/pkg/cmd/deregister/deregister_test.go @@ -1,387 +1,565 @@ package deregister import ( + "bytes" "context" - "fmt" - "net/http/httptest" - "os/user" + "errors" + "io" + "os" + "strings" "testing" nodev1connect "buf.build/gen/go/brevdev/devplane/connectrpc/go/devplaneapi/v1/devplaneapiv1connect" nodev1 "buf.build/gen/go/brevdev/devplane/protocolbuffers/go/devplaneapi/v1" "connectrpc.com/connect" + "github.com/spf13/cobra" + "github.com/stretchr/testify/require" "github.com/brevdev/brev-cli/pkg/cmd/register" "github.com/brevdev/brev-cli/pkg/entity" "github.com/brevdev/brev-cli/pkg/externalnode" - "github.com/brevdev/brev-cli/pkg/sudo" "github.com/brevdev/brev-cli/pkg/terminal" ) -type mockDeregisterStore struct { - user *entity.User - token string - err error +type leaveTestPlatform struct { + compatible bool + events *[]string } -func (m *mockDeregisterStore) GetCurrentUser() (*entity.User, error) { - if m.err != nil { - return nil, m.err - } - return m.user, nil +func (p *leaveTestPlatform) IsCompatible() bool { + recordLeaveEvent(p.events, "platform") + return p.compatible } -func (m *mockDeregisterStore) GetAccessToken() (string, error) { return m.token, nil } - -// fakeNodeService implements the server side of ExternalNodeService for testing. -type fakeNodeService struct { - nodev1connect.UnimplementedExternalNodeServiceHandler - removeNodeFn func(*nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) +type leaveTestStore struct { + events *[]string + currentUserCalls int + currentUserErr error } -func (f *fakeNodeService) RemoveNode(_ context.Context, req *connect.Request[nodev1.RemoveNodeRequest]) (*connect.Response[nodev1.RemoveNodeResponse], error) { - resp, err := f.removeNodeFn(req.Msg) - if err != nil { - return nil, err +func (s *leaveTestStore) GetCurrentUser() (*entity.User, error) { + s.currentUserCalls++ + recordLeaveEvent(s.events, "auth") + if s.currentUserErr != nil { + return nil, s.currentUserErr } - return connect.NewResponse(resp), nil + return &entity.User{ID: "user_current"}, nil } -// mockRegistrationStore satisfies register.RegistrationStore for deregister tests. -type mockRegistrationStore struct { - reg *register.DeviceRegistration +func (*leaveTestStore) GetAccessToken() (string, error) { return "token", nil } + +type leaveTestRegistrationStore struct { + events *[]string + reg *register.DeviceRegistration + loadErr error + deleteErr error + saveCalls int + deleteCalls int } -func (m *mockRegistrationStore) Save(reg *register.DeviceRegistration) error { - m.reg = reg +func (s *leaveTestRegistrationStore) Save(*register.DeviceRegistration) error { + s.saveCalls++ + recordLeaveEvent(s.events, "registration-save") return nil } -func (m *mockRegistrationStore) Load() (*register.DeviceRegistration, error) { - if m.reg == nil { - return nil, fmt.Errorf("no registration") +func (s *leaveTestRegistrationStore) Load() (*register.DeviceRegistration, error) { + recordLeaveEvent(s.events, "registration-load") + if s.loadErr != nil { + return nil, s.loadErr } - return m.reg, nil + return s.reg, nil } -func (m *mockRegistrationStore) Delete() error { - m.reg = nil +func (s *leaveTestRegistrationStore) Delete() error { + s.deleteCalls++ + recordLeaveEvent(s.events, "registration-delete") + if s.deleteErr != nil { + return s.deleteErr + } + s.reg = nil return nil } -func (m *mockRegistrationStore) Exists() (bool, error) { - return m.reg != nil, nil +func (s *leaveTestRegistrationStore) Exists() (bool, error) { return s.reg != nil, nil } + +type leaveTestConfirmer struct { + events *[]string + answer bool + calls int } -// mock types for deregisterDeps interfaces +func (c *leaveTestConfirmer) ConfirmYesNo(string) bool { + c.calls++ + recordLeaveEvent(c.events, "confirm") + return c.answer +} -type mockPlatform struct{ compatible bool } +type leaveTestGater struct { + events *[]string + calls int + reasons []string + err error +} -func (m mockPlatform) IsCompatible() bool { return m.compatible } +func (g *leaveTestGater) Gate(_ *terminal.Terminal, _ terminal.Confirmer, reason string, _ bool) error { + g.calls++ + g.reasons = append(g.reasons, reason) + recordLeaveEvent(g.events, "sudo") + return g.err +} -type mockSelector struct { - fn func(label string, items []string) string +type leaveTestNetBird struct { + events *[]string + calls int + err error } -func (m mockSelector) Select(label string, items []string) string { - return m.fn(label, items) +func (n *leaveTestNetBird) Uninstall() error { + n.calls++ + recordLeaveEvent(n.events, "netbird-uninstall") + return n.err } -type mockConfirmer struct{ confirm bool } +type leaveRecordingClient struct { + nodev1connect.ExternalNodeServiceClient + + events *[]string + listResponse *nodev1.ListNodesResponse + listErr error + returnNilList bool + removeErr error + listRequests []*nodev1.ListNodesRequest + removeRequests []*nodev1.RemoveNodeRequest + revokeCalls int +} -func (m mockConfirmer) ConfirmYesNo(_ string) bool { return m.confirm } +func (c *leaveRecordingClient) ListNodes(_ context.Context, req *connect.Request[nodev1.ListNodesRequest]) (*connect.Response[nodev1.ListNodesResponse], error) { + recordLeaveEvent(c.events, "list-nodes") + c.listRequests = append(c.listRequests, &nodev1.ListNodesRequest{OrganizationId: req.Msg.GetOrganizationId()}) + if c.listErr != nil { + return nil, c.listErr + } + if c.returnNilList { + return nil, nil + } + return connect.NewResponse(c.listResponse), nil +} -type mockNetBirdManager struct { - called bool - err error +func (c *leaveRecordingClient) RemoveNode(_ context.Context, req *connect.Request[nodev1.RemoveNodeRequest]) (*connect.Response[nodev1.RemoveNodeResponse], error) { + recordLeaveEvent(c.events, "remove-node") + c.removeRequests = append(c.removeRequests, &nodev1.RemoveNodeRequest{ExternalNodeId: req.Msg.GetExternalNodeId()}) + if c.removeErr != nil { + return nil, c.removeErr + } + return connect.NewResponse(&nodev1.RemoveNodeResponse{}), nil } -func (m *mockNetBirdManager) Install() error { return m.err } -func (m *mockNetBirdManager) Uninstall() error { m.called = true; return m.err } -func (m *mockNetBirdManager) EnsureRunning() error { return m.err } +func (c *leaveRecordingClient) RevokeNodeSSHAccess(context.Context, *connect.Request[nodev1.RevokeNodeSSHAccessRequest]) (*connect.Response[nodev1.RevokeNodeSSHAccessResponse], error) { + c.revokeCalls++ + recordLeaveEvent(c.events, "revoke-ssh") + return connect.NewResponse(&nodev1.RevokeNodeSSHAccessResponse{}), nil +} -type mockNodeClientFactory struct { - serverURL string +type leaveTestNodeClientFactory struct { + client nodev1connect.ExternalNodeServiceClient } -func (m mockNodeClientFactory) NewNodeClient(provider externalnode.TokenProvider, _ string) nodev1connect.ExternalNodeServiceClient { - return register.NewNodeServiceClient(provider, m.serverURL) +func (f leaveTestNodeClientFactory) NewNodeClient(externalnode.TokenProvider, string) nodev1connect.ExternalNodeServiceClient { + return f.client } -type mockSSHKeyRemover struct { - called bool - err error - removed []string +type leaveTestHarness struct { + events []string + store *leaveTestStore + registrations *leaveTestRegistrationStore + confirmer *leaveTestConfirmer + gater *leaveTestGater + netbird *leaveTestNetBird + client *leaveRecordingClient + deps leaveDeps } -func (m *mockSSHKeyRemover) RemoveBrevKeys(_ *user.User) ([]string, error) { - m.called = true - return m.removed, m.err +func newLeaveTestHarness() *leaveTestHarness { + h := &leaveTestHarness{} + reg := ®ister.DeviceRegistration{ + ExternalNodeID: "node_123", + DisplayName: "owned-node", + OrgID: "org_123", + OrgName: "owned-org", + } + node := &nodev1.ExternalNode{ExternalNodeId: reg.ExternalNodeID, Name: reg.DisplayName} + h.store = &leaveTestStore{events: &h.events} + h.registrations = &leaveTestRegistrationStore{events: &h.events, reg: reg} + h.confirmer = &leaveTestConfirmer{events: &h.events, answer: true} + h.gater = &leaveTestGater{events: &h.events} + h.netbird = &leaveTestNetBird{events: &h.events} + h.client = &leaveRecordingClient{ + events: &h.events, + listResponse: &nodev1.ListNodesResponse{Items: []*nodev1.ExternalNode{node}}, + } + h.deps = leaveDeps{ + platform: &leaveTestPlatform{compatible: true, events: &h.events}, + confirmer: h.confirmer, + gater: h.gater, + netbird: h.netbird, + nodeClients: leaveTestNodeClientFactory{client: h.client}, + registrationStore: h.registrations, + } + return h } -// testDeregisterDeps returns deps with all side-effects stubbed. The -// prompter defaults to confirming all prompts. -func testDeregisterDeps(t *testing.T, svc *fakeNodeService, regStore register.RegistrationStore) (deregisterDeps, *httptest.Server) { +func (h *leaveTestHarness) run(t *testing.T, approve bool) (stdout string, stderr string, err error) { t.Helper() + var warnings bytes.Buffer + stdout, err = captureLeaveStdout(t, func(term *terminal.Terminal) error { + return runLeave(context.Background(), term, &warnings, h.store, h.deps, approve) + }) + return stdout, warnings.String(), err +} - _, handler := nodev1connect.NewExternalNodeServiceHandler(svc) - server := httptest.NewServer(handler) - - return deregisterDeps{ - platform: mockPlatform{compatible: true}, - prompter: mockSelector{fn: func(_ string, items []string) string { - // Default: pick first item (Yes, ...) - if len(items) > 0 { - return items[0] - } - return "" - }}, - confirmer: mockConfirmer{confirm: true}, - gater: sudo.CachedGater{}, - netbird: &mockNetBirdManager{}, - nodeClients: mockNodeClientFactory{serverURL: server.URL}, - registrationStore: regStore, - sshKeys: &mockSSHKeyRemover{}, - }, server -} - -func Test_runDeregister_HappyPath(t *testing.T) { - regStore := &mockRegistrationStore{ - reg: ®ister.DeviceRegistration{ - ExternalNodeID: "unode_abc", - DisplayName: "My Spark", - OrgID: "org_123", - DeviceID: "dev-uuid", - }, - } - - store := &mockDeregisterStore{ - user: &entity.User{ID: "user_1"}, +func TestNewCmdLeave_CommandSurface(t *testing.T) { + cmd := NewCmdLeave(terminal.New(), &leaveTestStore{}) + require.Equal(t, "leave", cmd.Use) + require.Equal(t, []string{"deregister"}, cmd.Aliases) + require.NotNil(t, cmd.Args) + require.Contains(t, cmd.Annotations, "configuration") + require.NotNil(t, cmd.Flags().Lookup("approve")) +} - token: "tok", - } +func TestNewCmdDeregister_DeprecatedSourceCompatibility(t *testing.T) { + var store DeregisterStore = &leaveTestStore{} + cmd := NewCmdDeregister(terminal.New(), store) - var gotNodeID string - svc := &fakeNodeService{ - removeNodeFn: func(req *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { - gotNodeID = req.GetExternalNodeId() - return &nodev1.RemoveNodeResponse{}, nil - }, - } + require.Equal(t, "leave", cmd.Name()) + require.Equal(t, []string{"deregister"}, cmd.Aliases) +} - deps, server := testDeregisterDeps(t, svc, regStore) - defer server.Close() +func TestNewCmdLeave_DeregisterAliasWarnsOnExecution(t *testing.T) { + h := newLeaveTestHarness() + var stderr bytes.Buffer + _, err := captureLeaveStdout(t, func(term *terminal.Terminal) error { + root := &cobra.Command{Use: "brev"} + root.AddCommand(newCmdLeave(term, h.store, h.deps)) + root.SetArgs([]string{"deregister", "--approve"}) + root.SetOut(io.Discard) + root.SetErr(&stderr) + return root.Execute() + }) + require.NoError(t, err) + require.True(t, strings.HasPrefix(stderr.String(), "Warning: \"brev deregister\" is deprecated; use \"brev leave\" instead.\n"+ + "This command does not revoke SSH access grants; run \"brev disable-ssh\" before leaving if you want to revoke them.\n")) +} - term := terminal.New() - err := runDeregister(context.Background(), term, store, deps, false) - if err != nil { - t.Fatalf("runDeregister failed: %v", err) +func TestNewCmdLeave_HelpDoesNotWarn(t *testing.T) { + for _, name := range []string{"leave", "deregister"} { + t.Run(name, func(t *testing.T) { + h := newLeaveTestHarness() + var stderr bytes.Buffer + root := &cobra.Command{Use: "brev"} + root.AddCommand(newCmdLeave(terminal.New(), h.store, h.deps)) + root.SetArgs([]string{name, "--help"}) + root.SetOut(io.Discard) + root.SetErr(&stderr) + + require.NoError(t, root.Execute()) + require.NotContains(t, stderr.String(), "deprecated") + require.Empty(t, h.events) + }) } +} - if gotNodeID != "unode_abc" { - t.Errorf("expected node ID unode_abc, got %s", gotNodeID) - } +func TestNewCmdLeave_CanonicalInvocationDoesNotWarnAboutDeprecation(t *testing.T) { + h := newLeaveTestHarness() + var stderr bytes.Buffer + _, err := captureLeaveStdout(t, func(term *terminal.Terminal) error { + root := &cobra.Command{Use: "brev"} + root.AddCommand(newCmdLeave(term, h.store, h.deps)) + root.SetArgs([]string{"leave", "--approve"}) + root.SetOut(io.Discard) + root.SetErr(&stderr) + return root.Execute() + }) + require.NoError(t, err) + require.NotContains(t, stderr.String(), "deprecated") + require.Contains(t, stderr.String(), "may interrupt commands using Brev SSH") +} - // Registration should be deleted - exists, err := regStore.Exists() - if err != nil { - t.Fatalf("Exists error: %v", err) - } - if exists { - t.Error("expected registration to be deleted after deregister") +func TestNewCmdLeave_RejectsArguments(t *testing.T) { + for _, name := range []string{"leave", "deregister"} { + t.Run(name, func(t *testing.T) { + h := newLeaveTestHarness() + root := &cobra.Command{Use: "brev"} + root.AddCommand(newCmdLeave(terminal.New(), h.store, h.deps)) + root.SetArgs([]string{name, "unexpected"}) + root.SetOut(io.Discard) + root.SetErr(io.Discard) + require.Error(t, root.Execute()) + require.Empty(t, h.events) + }) } } -func Test_runDeregister_UserCancels(t *testing.T) { - regStore := &mockRegistrationStore{ - reg: ®ister.DeviceRegistration{ - ExternalNodeID: "unode_abc", - DisplayName: "My Spark", - OrgID: "org_123", - }, +func TestRunLeave_RemainingGrantsWarnButDoNotBlock(t *testing.T) { + h := newLeaveTestHarness() + h.client.listResponse.Items[0].SshAccess = []*nodev1.SSHAccess{ + {UserId: "user_1", LinuxUser: "ubuntu", PortId: "port_1"}, + nil, + {UserId: "user_2", LinuxUser: "ubuntu", PortId: "port_2"}, + {UserId: "user_3", LinuxUser: "alice", PortId: "port_3"}, } - store := &mockDeregisterStore{ - user: &entity.User{ID: "user_1"}, + _, stderr, err := h.run(t, false) + require.NoError(t, err) + require.Contains(t, stderr, "3 SSH grants across 2 Linux accounts") + require.Contains(t, stderr, `Leaving stops Brev-routed SSH but does not revoke these grants. Cancel and run "brev disable-ssh" first if you want them revoked.`) + require.Len(t, h.client.removeRequests, 1) +} - token: "tok", +func TestRunLeave_ApproveSkipsConfirmationButNotWarnings(t *testing.T) { + h := newLeaveTestHarness() + h.client.listResponse.Items[0].SshAccess = []*nodev1.SSHAccess{ + {UserId: "user_1", LinuxUser: "ubuntu", PortId: "port_1"}, + {UserId: "user_2", LinuxUser: "alice", PortId: "port_2"}, } + _, stderr, err := h.run(t, true) + require.NoError(t, err) + require.Zero(t, h.confirmer.calls) + require.Contains(t, stderr, "may interrupt commands using Brev SSH") + require.Contains(t, stderr, "2 SSH grants across 2 Linux accounts") + require.Contains(t, stderr, `run "brev disable-ssh" first`) + require.Equal(t, 1, h.gater.calls) +} - svc := &fakeNodeService{} - deps, server := testDeregisterDeps(t, svc, regStore) - defer server.Close() +func TestRunLeave_CancelStopsBeforeSudoAndMutation(t *testing.T) { + h := newLeaveTestHarness() + h.confirmer.answer = false + + stdout, _, err := h.run(t, false) + require.NoError(t, err) + require.Equal(t, []string{"platform", "registration-load", "auth", "list-nodes", "confirm"}, h.events) + require.Zero(t, h.gater.calls) + require.Empty(t, h.client.removeRequests) + require.Zero(t, h.netbird.calls) + require.Zero(t, h.registrations.deleteCalls) + require.NotContains(t, stdout, "Left the Brev network.") +} - deps.prompter = mockSelector{fn: func(_ string, _ []string) string { - return "No, cancel" - }} +func TestRunLeave_IncompatiblePlatformStopsBeforeLoadOrMutation(t *testing.T) { + h := newLeaveTestHarness() + h.deps.platform = &leaveTestPlatform{compatible: false, events: &h.events} + + stdout, _, err := h.run(t, false) + require.EqualError(t, err, "brev leave is only supported on Linux") + require.Equal(t, []string{"platform"}, h.events) + require.NotNil(t, h.registrations.reg) + require.Zero(t, h.confirmer.calls) + require.Zero(t, h.gater.calls) + require.Empty(t, h.client.removeRequests) + require.Zero(t, h.netbird.calls) + require.Zero(t, h.registrations.deleteCalls) + require.NotContains(t, stdout, "Left the Brev network.") +} - term := terminal.New() - err := runDeregister(context.Background(), term, store, deps, false) - if err != nil { - t.Fatalf("expected nil error on cancel, got: %v", err) - } +func TestRunLeave_AuthenticationFailureStopsBeforeLookupOrMutation(t *testing.T) { + h := newLeaveTestHarness() + authErr := errors.New("authentication failed") + h.store.currentUserErr = authErr + + stdout, _, err := h.run(t, false) + require.ErrorIs(t, err, authErr) + require.Equal(t, []string{"platform", "registration-load", "auth"}, h.events) + require.NotNil(t, h.registrations.reg) + require.Empty(t, h.client.listRequests) + require.Zero(t, h.confirmer.calls) + require.Zero(t, h.gater.calls) + require.Empty(t, h.client.removeRequests) + require.Zero(t, h.netbird.calls) + require.Zero(t, h.registrations.deleteCalls) + require.NotContains(t, stdout, "Left the Brev network.") +} - // Registration should still exist - exists, err := regStore.Exists() - if err != nil { - t.Fatalf("Exists error: %v", err) - } - if !exists { - t.Error("registration should still exist after cancel") - } +func TestRunLeave_SudoFailureStopsBeforeAuthoritativeMutation(t *testing.T) { + h := newLeaveTestHarness() + sudoErr := errors.New("sudo unavailable") + h.gater.err = sudoErr + + stdout, _, err := h.run(t, true) + require.ErrorIs(t, err, sudoErr) + require.Equal(t, []string{"platform", "registration-load", "auth", "list-nodes", "sudo"}, h.events) + require.NotNil(t, h.registrations.reg) + require.Empty(t, h.client.removeRequests) + require.Zero(t, h.netbird.calls) + require.Zero(t, h.registrations.deleteCalls) + require.NotContains(t, stdout, "Left the Brev network.") } -func Test_runDeregister_NotRegistered(t *testing.T) { - regStore := &mockRegistrationStore{} +func TestRunLeave_OrderIsRemoveNodeUninstallDeleteRegistration(t *testing.T) { + h := newLeaveTestHarness() + + stdout, _, err := h.run(t, false) + require.NoError(t, err) + require.Equal(t, []string{ + "platform", "registration-load", "auth", "list-nodes", "confirm", "sudo", + "remove-node", "netbird-uninstall", "registration-delete", + }, h.events) + require.Equal(t, []string{"Leave Brev network"}, h.gater.reasons) + require.Equal(t, []*nodev1.ListNodesRequest{{OrganizationId: "org_123"}}, h.client.listRequests) + require.Equal(t, []*nodev1.RemoveNodeRequest{{ExternalNodeId: "node_123"}}, h.client.removeRequests) + require.Contains(t, stdout, "Left the Brev network.") +} - store := &mockDeregisterStore{ - user: &entity.User{ID: "user_1"}, +func TestRunLeave_CompleteNodeListWithoutRegisteredIDAllowsAuthoritativeRemoveRetry(t *testing.T) { + h := newLeaveTestHarness() + h.client.listResponse = &nodev1.ListNodesResponse{Items: []*nodev1.ExternalNode{nil, {ExternalNodeId: "other"}}} - token: "tok", - } + _, stderr, err := h.run(t, true) + require.NoError(t, err) + require.Contains(t, stderr, "backend node is already absent") + require.Contains(t, stderr, "skipping SSH grant inspection") + require.Len(t, h.client.removeRequests, 1) +} - svc := &fakeNodeService{} - deps, server := testDeregisterDeps(t, svc, regStore) - defer server.Close() +func TestRunLeave_ListPermissionDeniedStopsBeforeConfirmationAndMutation(t *testing.T) { + h := newLeaveTestHarness() + h.client.listErr = connect.NewError(connect.CodePermissionDenied, errors.New("denied")) + assertLeaveLookupFailureStopsMutation(t, h) +} - term := terminal.New() - err := runDeregister(context.Background(), term, store, deps, false) - if err == nil { - t.Fatal("expected error when not registered") +func TestRunLeave_RegisteredIDAbsentFromIncompleteListStopsBeforeMutation(t *testing.T) { + h := newLeaveTestHarness() + h.client.listResponse = &nodev1.ListNodesResponse{ + Items: []*nodev1.ExternalNode{{ExternalNodeId: "other"}}, + NextPageToken: "next", } + assertLeaveLookupFailureStopsMutation(t, h) } -func Test_runDeregister_RemoveNodeFails(t *testing.T) { - regStore := &mockRegistrationStore{ - reg: ®ister.DeviceRegistration{ - ExternalNodeID: "unode_abc", - DisplayName: "My Spark", - OrgID: "org_123", - }, - } +func TestRunLeave_OtherLookupFailureStopsBeforeConfirmationAndMutation(t *testing.T) { + h := newLeaveTestHarness() + h.client.listErr = errors.New("backend unavailable") + assertLeaveLookupFailureStopsMutation(t, h) +} - store := &mockDeregisterStore{ - user: &entity.User{ID: "user_1"}, +func TestRunLeave_EmptyLookupResponseStopsBeforeConfirmationAndMutation(t *testing.T) { + h := newLeaveTestHarness() + h.client.returnNilList = true + assertLeaveLookupFailureStopsMutation(t, h) +} - token: "tok", - } +func TestRunLeave_EmptyLookupMessageStopsBeforeConfirmationAndMutation(t *testing.T) { + h := newLeaveTestHarness() + h.client.listResponse = nil + assertLeaveLookupFailureStopsMutation(t, h) +} - svc := &fakeNodeService{ - removeNodeFn: func(_ *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { - return nil, connect.NewError(connect.CodeInternal, nil) - }, - } +func TestRunLeave_RemoveNodeNotFoundIsAccepted(t *testing.T) { + h := newLeaveTestHarness() + h.client.removeErr = connect.NewError(connect.CodeNotFound, errors.New("already absent")) - deps, server := testDeregisterDeps(t, svc, regStore) - defer server.Close() + stdout, _, err := h.run(t, true) + require.NoError(t, err) + require.Equal(t, 1, h.netbird.calls) + require.Equal(t, 1, h.registrations.deleteCalls) + require.Contains(t, stdout, "Left the Brev network.") +} - term := terminal.New() - err := runDeregister(context.Background(), term, store, deps, false) - if err == nil { - t.Fatal("expected error when RemoveNode fails") - } +func TestRunLeave_RemoveNodeFailureStopsLocalTeardown(t *testing.T) { + h := newLeaveTestHarness() + removeErr := errors.New("remove failed") + h.client.removeErr = connect.NewError(connect.CodeInternal, removeErr) - // Registration should still exist (server-side removal failed) - exists, err := regStore.Exists() - if err != nil { - t.Fatalf("Exists error: %v", err) - } - if !exists { - t.Error("registration should still exist when RemoveNode fails") - } + stdout, _, err := h.run(t, true) + require.ErrorIs(t, err, removeErr) + require.Zero(t, h.netbird.calls) + require.Zero(t, h.registrations.deleteCalls) + require.NotContains(t, stdout, "Left the Brev network.") } -func Test_runDeregister_AlwaysUninstallsNetbird(t *testing.T) { - regStore := &mockRegistrationStore{ - reg: ®ister.DeviceRegistration{ - ExternalNodeID: "unode_abc", - DisplayName: "My Spark", - OrgID: "org_123", - }, - } +func TestRunLeave_NetBirdFailureReturnsErrorAndRetainsRegistration(t *testing.T) { + h := newLeaveTestHarness() + netbirdErr := errors.New("uninstall failed") + h.netbird.err = netbirdErr - store := &mockDeregisterStore{ - user: &entity.User{ID: "user_1"}, + stdout, _, err := h.run(t, true) + require.ErrorIs(t, err, netbirdErr) + require.Zero(t, h.registrations.deleteCalls) + require.NotNil(t, h.registrations.reg) + require.NotContains(t, stdout, "Left the Brev network.") +} - token: "tok", - } +func TestRunLeave_RegistrationDeleteFailureReturnsErrorAndNoSuccess(t *testing.T) { + h := newLeaveTestHarness() + deleteErr := errors.New("delete failed") + h.registrations.deleteErr = deleteErr - svc := &fakeNodeService{ - removeNodeFn: func(_ *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { - return &nodev1.RemoveNodeResponse{}, nil - }, - } + stdout, _, err := h.run(t, true) + require.ErrorIs(t, err, deleteErr) + require.Equal(t, 1, h.registrations.deleteCalls) + require.NotNil(t, h.registrations.reg) + require.NotContains(t, stdout, "Left the Brev network.") +} - netbird := &mockNetBirdManager{} - deps, server := testDeregisterDeps(t, svc, regStore) - defer server.Close() - deps.netbird = netbird +func TestRunLeave_NeverRevokesSSHOrSavesRegistration(t *testing.T) { + h := newLeaveTestHarness() + h.client.listResponse.Items[0].SshAccess = []*nodev1.SSHAccess{{UserId: "user_1", LinuxUser: "ubuntu", PortId: "port_1"}} - term := terminal.New() - err := runDeregister(context.Background(), term, store, deps, false) - if err != nil { - t.Fatalf("runDeregister failed: %v", err) - } + _, _, err := h.run(t, true) + require.NoError(t, err) + require.Zero(t, h.client.revokeCalls) + require.Zero(t, h.registrations.saveCalls) + require.NotContains(t, h.events, "revoke-ssh") +} - if !netbird.called { - t.Error("expected Brev tunnel uninstall to always be called during deregistration") - } +func TestRunLeave_RegistrationLoadFailureDoesNotAuthenticate(t *testing.T) { + h := newLeaveTestHarness() + h.registrations.loadErr = errors.New("registration missing") + + stdout, _, err := h.run(t, false) + require.Error(t, err) + require.Equal(t, []string{"platform", "registration-load"}, h.events) + require.Zero(t, h.store.currentUserCalls) + require.NotNil(t, h.registrations.reg) + require.Empty(t, h.client.removeRequests) + require.Zero(t, h.netbird.calls) + require.Zero(t, h.registrations.deleteCalls) + require.NotContains(t, stdout, "Left the Brev network.") } -func Test_runDeregister_RemoveBrevKeysHandling(t *testing.T) { - tests := []struct { - name string - sshKeys *mockSSHKeyRemover - wantCalled bool - }{ - {"CallsRemoveBrevKeys", &mockSSHKeyRemover{}, true}, - {"FailureIsNonFatal", &mockSSHKeyRemover{err: fmt.Errorf("permission denied")}, true}, - } +func assertLeaveLookupFailureStopsMutation(t *testing.T, h *leaveTestHarness) { + t.Helper() + stdout, _, err := h.run(t, false) + require.Error(t, err) + require.Contains(t, err.Error(), "inspect joined node before leaving") + require.Equal(t, []string{"platform", "registration-load", "auth", "list-nodes"}, h.events) + require.Zero(t, h.confirmer.calls) + require.Zero(t, h.gater.calls) + require.Empty(t, h.client.removeRequests) + require.Zero(t, h.netbird.calls) + require.Zero(t, h.registrations.deleteCalls) + require.NotNil(t, h.registrations.reg) + require.NotContains(t, stdout, "Left the Brev network.") +} - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - regStore := &mockRegistrationStore{ - reg: ®ister.DeviceRegistration{ - ExternalNodeID: "unode_abc", - DisplayName: "My Spark", - OrgID: "org_123", - }, - } - - store := &mockDeregisterStore{ - user: &entity.User{ID: "user_1"}, - - token: "tok", - } - - svc := &fakeNodeService{ - removeNodeFn: func(_ *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) { - return &nodev1.RemoveNodeResponse{}, nil - }, - } - - deps, server := testDeregisterDeps(t, svc, regStore) - defer server.Close() - deps.sshKeys = tt.sshKeys - - term := terminal.New() - err := runDeregister(context.Background(), term, store, deps, false) - if err != nil { - t.Fatalf("runDeregister failed: %v", err) - } - - if tt.sshKeys.called != tt.wantCalled { - t.Errorf("removeBrevKeys called = %v, want %v", tt.sshKeys.called, tt.wantCalled) - } - - // Registration should be cleaned up regardless of SSH key result. - exists, err := regStore.Exists() - if err != nil { - t.Fatalf("Exists error: %v", err) - } - if exists { - t.Error("expected registration to be deleted") - } - }) +func recordLeaveEvent(events *[]string, event string) { + if events != nil { + *events = append(*events, event) } } + +func captureLeaveStdout(t *testing.T, run func(*terminal.Terminal) error) (string, error) { + t.Helper() + reader, writer, err := os.Pipe() + require.NoError(t, err) + oldStdout := os.Stdout + os.Stdout = writer + term := terminal.New() + os.Stdout = oldStdout + + runErr := run(term) + require.NoError(t, writer.Close()) + output, readErr := io.ReadAll(reader) + require.NoError(t, readErr) + require.NoError(t, reader.Close()) + return string(output), runErr +} diff --git a/pkg/cmd/disablessh/disablessh.go b/pkg/cmd/disablessh/disablessh.go new file mode 100644 index 000000000..7fe03ed15 --- /dev/null +++ b/pkg/cmd/disablessh/disablessh.go @@ -0,0 +1,177 @@ +// Package disablessh provides the node-wide brev disable-ssh command. +package disablessh + +import ( + "context" + "fmt" + "io" + + nodev1connect "buf.build/gen/go/brevdev/devplane/connectrpc/go/devplaneapi/v1/devplaneapiv1connect" + nodev1 "buf.build/gen/go/brevdev/devplane/protocolbuffers/go/devplaneapi/v1" + "connectrpc.com/connect" + + "github.com/brevdev/brev-cli/pkg/cmd/register" + "github.com/brevdev/brev-cli/pkg/config" + "github.com/brevdev/brev-cli/pkg/entity" + breverrors "github.com/brevdev/brev-cli/pkg/errors" + "github.com/brevdev/brev-cli/pkg/externalnode" + "github.com/brevdev/brev-cli/pkg/terminal" + + "github.com/spf13/cobra" +) + +// DisableSSHStore defines the authenticated store methods needed by disable-ssh. +type DisableSSHStore interface { + GetCurrentUser() (*entity.User, error) + GetAccessToken() (string, error) +} + +type disableSSHDeps struct { + confirmer terminal.Confirmer + nodeClients externalnode.NodeClientFactory + registrationStore register.RegistrationStore +} + +func defaultDisableSSHDeps() disableSSHDeps { + return disableSSHDeps{ + confirmer: register.TerminalPrompter{}, + nodeClients: register.DefaultNodeClientFactory{}, + registrationStore: register.NewFileRegistrationStore(), + } +} + +// NewCmdDisableSSH creates the canonical node-wide disable-ssh command. +func NewCmdDisableSSH(t *terminal.Terminal, store DisableSSHStore) *cobra.Command { + return newCmdDisableSSH(t, store, defaultDisableSSHDeps()) +} + +func newCmdDisableSSH(t *terminal.Terminal, store DisableSSHStore, deps disableSSHDeps) *cobra.Command { + var approveFlag bool + cmd := &cobra.Command{ + Annotations: map[string]string{"configuration": ""}, + Use: "disable-ssh", + DisableFlagsInUseLine: true, + Short: "Revoke all Brev SSH access grants on this node", + Long: "Revoke every Brev SSH access grant on this joined node without changing Brev network membership or the SSH daemon.", + Example: " brev disable-ssh\n brev disable-ssh --approve", + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, _ []string) error { + return runDisableSSH(cmd.Context(), t, cmd.ErrOrStderr(), store, deps, approveFlag) + }, + } + cmd.Flags().BoolVar(&approveFlag, "approve", false, "skip confirmation prompt (assume yes)") + return cmd +} + +func runDisableSSH( + ctx context.Context, + t *terminal.Terminal, + warnings io.Writer, + store DisableSSHStore, + deps disableSSHDeps, + skipConfirm bool, +) error { //nolint:funlen // Keep the node-wide confirmation and revocation flow linear. + exists, err := deps.registrationStore.Exists() + if err != nil { + return fmt.Errorf("check joined-device registration: %w", err) + } + if !exists { + return breverrors.New(`This machine has not joined a Brev network; run "brev join" first.`) + } + + reg, err := deps.registrationStore.Load() + if err != nil { + return fmt.Errorf("read joined-device registration: %w", err) + } + currentUser, err := store.GetCurrentUser() + if err != nil { + return breverrors.WrapAndTrace(err) + } + if currentUser == nil || currentUser.ID == "" { + return fmt.Errorf("get current Brev user: missing user ID") + } + + node, err := register.FetchRegisteredNode(ctx, deps.nodeClients, store, reg) + if err != nil { + return fmt.Errorf("disable SSH failed: %w", err) + } + accesses := snapshotSSHAccessForRevocation(node.GetSshAccess(), currentUser.ID) + + t.Vprint("") + t.Vprint(t.White("════════════════════════════════════════════")) + t.Vprint(t.White(" Disabling Brev SSH access")) + t.Vprint(t.White("════════════════════════════════════════════")) + t.Vprint("") + t.Vprintf(" Node: %s (%s)\n", node.GetName(), node.GetExternalNodeId()) + t.Vprintf(" SSH grants: %d\n", len(accesses)) + t.Vprint("") + if len(accesses) == 0 { + t.Vprint(t.Green("No SSH access grants to revoke.")) + return nil + } + + if warnings == nil { + warnings = io.Discard + } + _, _ = fmt.Fprintln(warnings, "Warning: this is a node-wide operation that revokes all Brev SSH access grants on this node.") + _, _ = fmt.Fprintln(warnings, "Warning: active SSH sessions are not forcibly terminated.") + + if !skipConfirm && !deps.confirmer.ConfirmYesNo("Disable all Brev-managed SSH access on this node?") { + t.Vprint("Disable SSH canceled.") + return nil + } + + client := deps.nodeClients.NewNodeClient(store, config.GlobalConfig.GetBrevPublicAPIURL()) + if err := revokeSSHAccesses(ctx, client, reg.ExternalNodeID, accesses); err != nil { + return err + } + + t.Vprintf("%s SSH access disabled. Grants revoked: %d.\n", t.Green(" ✓"), len(accesses)) + return nil +} + +func revokeSSHAccesses( + ctx context.Context, + client nodev1connect.ExternalNodeServiceClient, + nodeID string, + accesses []*nodev1.SSHAccess, +) error { + var revokeErrs []error + for _, access := range accesses { + _, err := client.RevokeNodeSSHAccess(ctx, connect.NewRequest(&nodev1.RevokeNodeSSHAccessRequest{ + ExternalNodeId: nodeID, + PortId: access.GetPortId(), + UserId: access.GetUserId(), + LinuxUser: access.GetLinuxUser(), + })) + if err != nil { + revokeErrs = append(revokeErrs, fmt.Errorf( + "revoke SSH access for user %q, Linux account %q, port %q: %w", + access.GetUserId(), + access.GetLinuxUser(), + access.GetPortId(), + err, + )) + } + } + if err := breverrors.Join(revokeErrs...); err != nil { + return fmt.Errorf("failed to revoke one or more SSH access grants: %w", err) + } + return nil +} + +func snapshotSSHAccessForRevocation(accesses []*nodev1.SSHAccess, currentUserID string) []*nodev1.SSHAccess { + snapshot := make([]*nodev1.SSHAccess, 0, len(accesses)) + currentUserAccesses := make([]*nodev1.SSHAccess, 0, len(accesses)) + for _, access := range accesses { + if access == nil { + continue + } + if access.GetUserId() == currentUserID { + currentUserAccesses = append(currentUserAccesses, access) + continue + } + snapshot = append(snapshot, access) + } + return append(snapshot, currentUserAccesses...) +} diff --git a/pkg/cmd/disablessh/disablessh_test.go b/pkg/cmd/disablessh/disablessh_test.go new file mode 100644 index 000000000..280c17d74 --- /dev/null +++ b/pkg/cmd/disablessh/disablessh_test.go @@ -0,0 +1,370 @@ +package disablessh + +import ( + "bytes" + "context" + "errors" + "fmt" + "io" + "os" + "testing" + + nodev1connect "buf.build/gen/go/brevdev/devplane/connectrpc/go/devplaneapi/v1/devplaneapiv1connect" + nodev1 "buf.build/gen/go/brevdev/devplane/protocolbuffers/go/devplaneapi/v1" + "connectrpc.com/connect" + "github.com/stretchr/testify/require" + + "github.com/brevdev/brev-cli/pkg/cmd/register" + "github.com/brevdev/brev-cli/pkg/entity" + "github.com/brevdev/brev-cli/pkg/externalnode" + "github.com/brevdev/brev-cli/pkg/terminal" +) + +type disableSSHTestStore struct { + currentUser *entity.User + currentUserErr error + currentUserCalls int +} + +func (s *disableSSHTestStore) GetCurrentUser() (*entity.User, error) { + s.currentUserCalls++ + return s.currentUser, s.currentUserErr +} + +func (*disableSSHTestStore) GetAccessToken() (string, error) { + return "token", nil +} + +type disableSSHTestRegistrationStore struct { + exists bool + existsErr error + loadErr error + reg *register.DeviceRegistration + saveCalls int + deleteCalls int +} + +func (s *disableSSHTestRegistrationStore) Exists() (bool, error) { + return s.exists, s.existsErr +} + +func (s *disableSSHTestRegistrationStore) Load() (*register.DeviceRegistration, error) { + return s.reg, s.loadErr +} + +func (s *disableSSHTestRegistrationStore) Save(*register.DeviceRegistration) error { + s.saveCalls++ + return nil +} + +func (s *disableSSHTestRegistrationStore) Delete() error { + s.deleteCalls++ + return nil +} + +type disableSSHTestConfirmer struct { + answer bool + calls int +} + +func (c *disableSSHTestConfirmer) ConfirmYesNo(string) bool { + c.calls++ + return c.answer +} + +type disableSSHRecordingClient struct { + nodev1connect.ExternalNodeServiceClient + + node *nodev1.ExternalNode + getErr error + getNodeCalls int + revokeErrors map[int]error + revokeRequests []*nodev1.RevokeNodeSSHAccessRequest + + addNodeCalls int + removeNodeCalls int + closePortCalls int +} + +func (c *disableSSHRecordingClient) GetNode(_ context.Context, req *connect.Request[nodev1.GetNodeRequest]) (*connect.Response[nodev1.GetNodeResponse], error) { + c.getNodeCalls++ + if c.getErr != nil { + return nil, c.getErr + } + if req.Msg.GetExternalNodeId() != "node_123" || req.Msg.GetOrganizationId() != "org_123" { + return nil, fmt.Errorf("unexpected GetNode request: %+v", req.Msg) + } + return connect.NewResponse(&nodev1.GetNodeResponse{ExternalNode: c.node}), nil +} + +func (c *disableSSHRecordingClient) RevokeNodeSSHAccess(_ context.Context, req *connect.Request[nodev1.RevokeNodeSSHAccessRequest]) (*connect.Response[nodev1.RevokeNodeSSHAccessResponse], error) { + callIndex := len(c.revokeRequests) + c.revokeRequests = append(c.revokeRequests, cloneRevokeRequest(req.Msg)) + if err := c.revokeErrors[callIndex]; err != nil { + return nil, err + } + return connect.NewResponse(&nodev1.RevokeNodeSSHAccessResponse{}), nil +} + +func (c *disableSSHRecordingClient) AddNode(context.Context, *connect.Request[nodev1.AddNodeRequest]) (*connect.Response[nodev1.AddNodeResponse], error) { + c.addNodeCalls++ + return connect.NewResponse(&nodev1.AddNodeResponse{}), nil +} + +func (c *disableSSHRecordingClient) RemoveNode(context.Context, *connect.Request[nodev1.RemoveNodeRequest]) (*connect.Response[nodev1.RemoveNodeResponse], error) { + c.removeNodeCalls++ + return connect.NewResponse(&nodev1.RemoveNodeResponse{}), nil +} + +func (c *disableSSHRecordingClient) ClosePort(context.Context, *connect.Request[nodev1.ClosePortRequest]) (*connect.Response[nodev1.ClosePortResponse], error) { + c.closePortCalls++ + return connect.NewResponse(&nodev1.ClosePortResponse{}), nil +} + +type disableSSHTestNodeClientFactory struct { + client nodev1connect.ExternalNodeServiceClient +} + +func (f disableSSHTestNodeClientFactory) NewNodeClient(externalnode.TokenProvider, string) nodev1connect.ExternalNodeServiceClient { + return f.client +} + +type disableSSHTestHarness struct { + store *disableSSHTestStore + registrations *disableSSHTestRegistrationStore + confirmer *disableSSHTestConfirmer + client *disableSSHRecordingClient + deps disableSSHDeps +} + +func newDisableSSHTestHarness(accesses ...*nodev1.SSHAccess) *disableSSHTestHarness { + h := &disableSSHTestHarness{ + store: &disableSSHTestStore{currentUser: &entity.User{ID: "user_current"}}, + registrations: &disableSSHTestRegistrationStore{ + exists: true, + reg: ®ister.DeviceRegistration{ + ExternalNodeID: "node_123", + DisplayName: "owned-node", + OrgID: "org_123", + OrgName: "owned-org", + }, + }, + confirmer: &disableSSHTestConfirmer{answer: true}, + client: &disableSSHRecordingClient{ + node: &nodev1.ExternalNode{ + ExternalNodeId: "node_123", + Name: "owned-node", + SshAccess: accesses, + }, + revokeErrors: make(map[int]error), + }, + } + h.deps = disableSSHDeps{ + confirmer: h.confirmer, + nodeClients: disableSSHTestNodeClientFactory{client: h.client}, + registrationStore: h.registrations, + } + return h +} + +func (h *disableSSHTestHarness) run(t *testing.T, skipConfirm bool) (stdout string, stderr string, err error) { + t.Helper() + var warnings bytes.Buffer + stdout, err = captureDisableSSHStdout(t, func(term *terminal.Terminal) error { + return runDisableSSH(context.Background(), term, &warnings, h.store, h.deps, skipConfirm) + }) + return stdout, warnings.String(), err +} + +func TestNewCmdDisableSSH_CommandSurface(t *testing.T) { + cmd := NewCmdDisableSSH(terminal.New(), &disableSSHTestStore{}) + require.Equal(t, "disable-ssh", cmd.Use) + require.Equal(t, "Revoke all Brev SSH access grants on this node", cmd.Short) + require.NotNil(t, cmd.Args) + require.Contains(t, cmd.Annotations, "configuration") + require.Empty(t, cmd.Aliases) + require.NotNil(t, cmd.Flags().Lookup("approve")) +} + +func TestNewCmdDisableSSH_RejectsArguments(t *testing.T) { + h := newDisableSSHTestHarness() + cmd := newCmdDisableSSH(terminal.New(), h.store, h.deps) + cmd.SetArgs([]string{"unexpected"}) + cmd.SetOut(io.Discard) + cmd.SetErr(io.Discard) + + err := cmd.Execute() + require.Error(t, err) + require.Empty(t, h.client.revokeRequests) +} + +func TestRunDisableSSH_MissingRegistrationDoesNotAuthenticateOrCallRPC(t *testing.T) { + h := newDisableSSHTestHarness() + h.registrations.exists = false + + _, _, err := h.run(t, false) + require.EqualError(t, err, `This machine has not joined a Brev network; run "brev join" first.`) + require.Zero(t, h.store.currentUserCalls) + require.Zero(t, h.client.getNodeCalls) + require.Empty(t, h.client.revokeRequests) +} + +func TestRunDisableSSH_RequiresCurrentUserIDBeforeLoadingGrants(t *testing.T) { + authErr := errors.New("authentication failed") + tests := []struct { + name string + currentUser *entity.User + currentErr error + }{ + {name: "authentication failure", currentErr: authErr}, + {name: "missing user"}, + {name: "missing user ID", currentUser: &entity.User{}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + h := newDisableSSHTestHarness(testSSHAccess("user_current", "ubuntu", "port_1")) + h.store.currentUser = tt.currentUser + h.store.currentUserErr = tt.currentErr + + _, _, err := h.run(t, true) + require.Error(t, err) + if tt.currentErr != nil { + require.ErrorIs(t, err, tt.currentErr) + } else { + require.Contains(t, err.Error(), "missing user ID") + } + require.Zero(t, h.client.getNodeCalls) + require.Zero(t, h.confirmer.calls) + require.Empty(t, h.client.revokeRequests) + }) + } +} + +func TestRunDisableSSH_NoGrantsReportsAlreadyDisabledWithoutPrompting(t *testing.T) { + h := newDisableSSHTestHarness() + + stdout, stderr, err := h.run(t, false) + require.NoError(t, err) + require.Contains(t, stdout, "No SSH access grants to revoke.") + require.NotContains(t, stdout, "keys removed") + require.Empty(t, stderr) + require.Zero(t, h.confirmer.calls) + require.Empty(t, h.client.revokeRequests) +} + +func TestRunDisableSSH_CancelDoesNotRevoke(t *testing.T) { + h := newDisableSSHTestHarness(testSSHAccess("user_collaborator", "ubuntu", "port_1")) + h.confirmer.answer = false + + stdout, stderr, err := h.run(t, false) + require.NoError(t, err) + require.Contains(t, stdout, "Disable SSH canceled.") + require.Contains(t, stderr, "node-wide operation") + require.Equal(t, 1, h.confirmer.calls) + require.Empty(t, h.client.revokeRequests) +} + +func TestRunDisableSSH_ApproveRevokesWithoutPrompting(t *testing.T) { + h := newDisableSSHTestHarness(testSSHAccess("user_collaborator", "ubuntu", "port_1")) + + stdout, stderr, err := h.run(t, true) + require.NoError(t, err) + require.Zero(t, h.confirmer.calls) + require.Contains(t, stderr, "active SSH sessions are not forcibly terminated") + require.Contains(t, stdout, "SSH access disabled. Grants revoked: 1.") + require.Len(t, h.client.revokeRequests, 1) +} + +func TestRunDisableSSH_RevokesEveryExactTupleWithCurrentUsersAccessLast(t *testing.T) { + h := newDisableSSHTestHarness( + testSSHAccess("user_current", "ubuntu", "port_self_1"), + testSSHAccess("user_collaborator_1", "alice", "port_collaborator_1"), + nil, + testSSHAccess("user_current", "root", "port_self_2"), + testSSHAccess("user_collaborator_2", "carol", "port_collaborator_2"), + ) + + _, _, err := h.run(t, true) + require.NoError(t, err) + require.Equal(t, []*nodev1.RevokeNodeSSHAccessRequest{ + {ExternalNodeId: "node_123", UserId: "user_collaborator_1", LinuxUser: "alice", PortId: "port_collaborator_1"}, + {ExternalNodeId: "node_123", UserId: "user_collaborator_2", LinuxUser: "carol", PortId: "port_collaborator_2"}, + {ExternalNodeId: "node_123", UserId: "user_current", LinuxUser: "ubuntu", PortId: "port_self_1"}, + {ExternalNodeId: "node_123", UserId: "user_current", LinuxUser: "root", PortId: "port_self_2"}, + }, h.client.revokeRequests) +} + +func TestRunDisableSSH_ContinuesAfterFailuresAndReturnsEveryCause(t *testing.T) { + firstErr := errors.New("first revoke failed") + secondErr := errors.New("second revoke failed") + h := newDisableSSHTestHarness( + testSSHAccess("user_1", "ubuntu", "port_1"), + testSSHAccess("user_2", "alice", "port_2"), + testSSHAccess("user_current", "root", "port_self"), + ) + h.client.revokeErrors[0] = firstErr + h.client.revokeErrors[1] = secondErr + + _, _, err := h.run(t, true) + require.ErrorIs(t, err, firstErr) + require.ErrorIs(t, err, secondErr) + require.Contains(t, err.Error(), "failed to revoke one or more SSH access grants") + for _, text := range []string{"user_1", "ubuntu", "port_1", "user_2", "alice", "port_2"} { + require.Contains(t, err.Error(), text) + } + require.Len(t, h.client.revokeRequests, 3) +} + +func TestRunDisableSSH_DoesNotChangeMembershipPortsOrRegistration(t *testing.T) { + h := newDisableSSHTestHarness(testSSHAccess("user_collaborator", "ubuntu", "port_1")) + + _, _, err := h.run(t, true) + require.NoError(t, err) + require.Zero(t, h.client.addNodeCalls) + require.Zero(t, h.client.removeNodeCalls) + require.Zero(t, h.client.closePortCalls) + require.Zero(t, h.registrations.saveCalls) + require.Zero(t, h.registrations.deleteCalls) +} + +func TestRunDisableSSH_BackendNodeFailureStopsBeforeConfirmation(t *testing.T) { + h := newDisableSSHTestHarness() + h.client.getErr = errors.New("backend unavailable") + + _, _, err := h.run(t, false) + require.Error(t, err) + require.Contains(t, err.Error(), "disable SSH failed") + require.Zero(t, h.confirmer.calls) + require.Empty(t, h.client.revokeRequests) +} + +func cloneRevokeRequest(req *nodev1.RevokeNodeSSHAccessRequest) *nodev1.RevokeNodeSSHAccessRequest { + return &nodev1.RevokeNodeSSHAccessRequest{ + ExternalNodeId: req.GetExternalNodeId(), + PortId: req.GetPortId(), + UserId: req.GetUserId(), + LinuxUser: req.GetLinuxUser(), + } +} + +func testSSHAccess(userID, linuxUser, portID string) *nodev1.SSHAccess { + return &nodev1.SSHAccess{UserId: userID, LinuxUser: linuxUser, PortId: portID} +} + +func captureDisableSSHStdout(t *testing.T, run func(*terminal.Terminal) error) (string, error) { + t.Helper() + reader, writer, err := os.Pipe() + require.NoError(t, err) + oldStdout := os.Stdout + os.Stdout = writer + term := terminal.New() + os.Stdout = oldStdout + + runErr := run(term) + require.NoError(t, writer.Close()) + output, readErr := io.ReadAll(reader) + require.NoError(t, readErr) + require.NoError(t, reader.Close()) + return string(output), runErr +} diff --git a/pkg/cmd/enablessh/enablessh.go b/pkg/cmd/enablessh/enablessh.go index 9788b0e6f..24065857f 100644 --- a/pkg/cmd/enablessh/enablessh.go +++ b/pkg/cmd/enablessh/enablessh.go @@ -9,10 +9,8 @@ import ( "os/user" nodev1 "buf.build/gen/go/brevdev/devplane/protocolbuffers/go/devplaneapi/v1" - "connectrpc.com/connect" "github.com/brevdev/brev-cli/pkg/cmd/register" - "github.com/brevdev/brev-cli/pkg/config" "github.com/brevdev/brev-cli/pkg/entity" breverrors "github.com/brevdev/brev-cli/pkg/errors" "github.com/brevdev/brev-cli/pkg/externalnode" @@ -27,21 +25,44 @@ type EnableSSHStore interface { GetAccessToken() (string, error) } +type sshAccessProvisioner interface { + Provision( + context.Context, + *terminal.Terminal, + externalnode.TokenProvider, + *register.DeviceRegistration, + *entity.User, + *nodev1.ExternalNode, + ) error +} + // enableSSHDeps bundles the side-effecting dependencies of runEnableSSH so they // can be replaced in tests. type enableSSHDeps struct { platform externalnode.PlatformChecker nodeClients externalnode.NodeClientFactory registrationStore register.RegistrationStore - prompter terminal.Selector + tunnel register.NetBirdConnector + provisioner sshAccessProvisioner +} + +type defaultSSHAccessProvisioner struct { + prompter terminal.Selector + nodeClients externalnode.NodeClientFactory } func defaultEnableSSHDeps() enableSSHDeps { + prompter := register.TerminalPrompter{} + nodeClients := register.DefaultNodeClientFactory{} return enableSSHDeps{ platform: register.LinuxPlatform{}, - nodeClients: register.DefaultNodeClientFactory{}, + nodeClients: nodeClients, registrationStore: register.NewFileRegistrationStore(), - prompter: register.TerminalPrompter{}, + tunnel: register.Netbird{}, + provisioner: defaultSSHAccessProvisioner{ + prompter: prompter, + nodeClients: nodeClients, + }, } } @@ -50,9 +71,10 @@ func NewCmdEnableSSH(t *terminal.Terminal, store EnableSSHStore) *cobra.Command Annotations: map[string]string{"configuration": ""}, Use: "enable-ssh", DisableFlagsInUseLine: true, - Short: "Enable SSH access to this registered device", - Long: "Enable SSH access to this registered device for the current Brev user.", + Short: "Enable SSH access to this joined node", + Long: "Enable SSH access to this joined node for the current Brev user.", Example: " brev enable-ssh", + Args: cobra.NoArgs, RunE: func(cmd *cobra.Command, args []string) error { return runEnableSSH(cmd.Context(), t, store, defaultEnableSSHDeps()) }, @@ -66,9 +88,17 @@ func runEnableSSH(ctx context.Context, t *terminal.Terminal, s EnableSSHStore, d return fmt.Errorf("brev enable-ssh is only supported on Linux") } + exists, err := deps.registrationStore.Exists() + if err != nil { + return fmt.Errorf("check joined-device registration: %w", err) + } + if !exists { + return breverrors.New(`This machine has not joined a Brev network; run "brev join" first.`) + } + reg, err := deps.registrationStore.Load() if err != nil { - return fmt.Errorf("failed to read registration file: %w", err) + return fmt.Errorf("read joined-device registration: %w", err) } brevUser, err := s.GetCurrentUser() @@ -76,18 +106,30 @@ func runEnableSSH(ctx context.Context, t *terminal.Terminal, s EnableSSHStore, d return breverrors.WrapAndTrace(err) } - return enableSSH(ctx, t, deps, s, reg, brevUser) + node, err := register.FetchRegisteredNode(ctx, deps.nodeClients, s, reg) + if err != nil { + return fmt.Errorf("enable SSH failed: %w", err) + } + if err := deps.tunnel.EnsureConnected(ctx); err != nil { + return fmt.Errorf("enable SSH requires a connected Brev tunnel: %w", err) + } + if err := deps.provisioner.Provision(ctx, t, s, reg, brevUser, node); err != nil { + return fmt.Errorf("enable SSH failed: %w", err) + } + + t.Vprint(t.Green(fmt.Sprintf("SSH access enabled. You can now SSH to this device via: brev shell %s", reg.DisplayName))) + return nil } -// enableSSH grants SSH access to the given node for the current Brev user. +// Provision grants SSH access to the joined node for the current Brev user. // This is the "reflexive grant" — granting yourself SSH access to the device. -func enableSSH( +func (p defaultSSHAccessProvisioner) Provision( ctx context.Context, t *terminal.Terminal, - deps enableSSHDeps, tokenProvider externalnode.TokenProvider, reg *register.DeviceRegistration, brevUser *entity.User, + node *nodev1.ExternalNode, ) error { linuxUser, err := user.Current() if err != nil { @@ -105,41 +147,18 @@ func enableSSH( t.Vprintf(" Linux user: %s\n", linuxUsername) t.Vprint("") - node, err := fetchRegisteredNode(ctx, deps, tokenProvider, reg) - if err != nil { - return fmt.Errorf("enable SSH failed: %w", err) - } - - brevPortID, err := register.ResolveSSHAccessPort(ctx, t, deps.prompter, deps.nodeClients, tokenProvider, reg, node) + brevPortID, err := register.ResolveSSHAccessPort(ctx, t, p.prompter, p.nodeClients, tokenProvider, reg, node) if err != nil { - return fmt.Errorf("enable SSH failed: %w", err) + return err //nolint:wrapcheck // ResolveSSHAccessPort returns operation-specific user guidance. } - if err := register.SetupAndRegisterNodeSSHAccess(ctx, t, deps.nodeClients, tokenProvider, reg, brevUser, linuxUsername, brevPortID); err != nil { - return fmt.Errorf("enable SSH failed: %w", err) + if err := register.SetupAndRegisterNodeSSHAccess(ctx, t, p.nodeClients, tokenProvider, reg, brevUser, linuxUsername, brevPortID); err != nil { + return err //nolint:wrapcheck // SetupAndRegisterNodeSSHAccess supplies provisioning context. } - t.Vprint(t.Green(fmt.Sprintf("SSH access enabled. You can now SSH to this device via: brev shell %s", reg.DisplayName))) return nil } -func fetchRegisteredNode( - ctx context.Context, - deps enableSSHDeps, - tokenProvider externalnode.TokenProvider, - reg *register.DeviceRegistration, -) (*nodev1.ExternalNode, error) { - client := deps.nodeClients.NewNodeClient(tokenProvider, config.GlobalConfig.GetBrevPublicAPIURL()) - resp, err := client.GetNode(ctx, connect.NewRequest(&nodev1.GetNodeRequest{ - ExternalNodeId: reg.ExternalNodeID, - OrganizationId: reg.OrgID, - })) - if err != nil { - return nil, fmt.Errorf("error retrieving node: %w", err) - } - return resp.Msg.GetExternalNode(), nil -} - // checkSSHDaemon prints a warning if neither "ssh" nor "sshd" systemd services // appear to be active. It never returns an error — it is best-effort. func checkSSHDaemon(t *terminal.Terminal) { diff --git a/pkg/cmd/enablessh/enablessh_test.go b/pkg/cmd/enablessh/enablessh_test.go index 7df94144d..1c90df270 100644 --- a/pkg/cmd/enablessh/enablessh_test.go +++ b/pkg/cmd/enablessh/enablessh_test.go @@ -2,280 +2,326 @@ package enablessh import ( "context" + "errors" "net/http/httptest" - "os" - "os/user" - "path/filepath" - "strings" "testing" nodev1connect "buf.build/gen/go/brevdev/devplane/connectrpc/go/devplaneapi/v1/devplaneapiv1connect" nodev1 "buf.build/gen/go/brevdev/devplane/protocolbuffers/go/devplaneapi/v1" "connectrpc.com/connect" + "github.com/stretchr/testify/require" "github.com/brevdev/brev-cli/pkg/cmd/register" + "github.com/brevdev/brev-cli/pkg/entity" "github.com/brevdev/brev-cli/pkg/externalnode" + "github.com/brevdev/brev-cli/pkg/terminal" ) -// tempUser returns a *user.User whose HomeDir points to a temporary directory. -func tempUser(t *testing.T) *user.User { - t.Helper() - return &user.User{HomeDir: t.TempDir()} +type mockNodeClientFactory struct{ serverURL string } + +func (m mockNodeClientFactory) NewNodeClient(provider externalnode.TokenProvider, _ string) nodev1connect.ExternalNodeServiceClient { + return register.NewNodeServiceClient(provider, m.serverURL) } -// readAuthorizedKeys is a test helper that reads ~/.ssh/authorized_keys. -func readAuthorizedKeys(t *testing.T, u *user.User) string { - t.Helper() - data, err := os.ReadFile(filepath.Join(u.HomeDir, ".ssh", "authorized_keys")) - if err != nil { - t.Fatalf("reading authorized_keys: %v", err) - } - return string(data) +type mockEnableSSHStore struct { + token string + user *entity.User + err error } -// --- RemoveBrevAuthorizedKeys --- +func (m *mockEnableSSHStore) GetCurrentUser() (*entity.User, error) { return m.user, m.err } +func (m *mockEnableSSHStore) GetAccessToken() (string, error) { return m.token, nil } -func Test_RemoveBrevAuthorizedKeys_RemovesTaggedKeys(t *testing.T) { - u := tempUser(t) - sshDir := filepath.Join(u.HomeDir, ".ssh") - if err := os.MkdirAll(sshDir, 0o700); err != nil { - t.Fatal(err) - } +// fakeNodeService implements the server side of ExternalNodeService for testing. +type fakeNodeService struct { + nodev1connect.UnimplementedExternalNodeServiceHandler + getNodeFn func(*nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) + order *[]string + addNodeCalls int +} - content := strings.Join([]string{ - "ssh-rsa EXISTING user@host", - "ssh-rsa BREVKEY1 " + register.DevplaneAuthorizedKeysComment("p1", "u1"), - "ssh-ed25519 OTHERKEY admin@server", - "ssh-rsa BREVKEY2 " + register.DevplaneAuthorizedKeysComment("p2", "u2"), - "", - }, "\n") - if err := os.WriteFile(filepath.Join(sshDir, "authorized_keys"), []byte(content), 0o600); err != nil { - t.Fatal(err) +func (f *fakeNodeService) GetNode(_ context.Context, req *connect.Request[nodev1.GetNodeRequest]) (*connect.Response[nodev1.GetNodeResponse], error) { + if f.order != nil { + *f.order = append(*f.order, "node") } - - removed, err := register.RemoveBrevAuthorizedKeys(u) + resp, err := f.getNodeFn(req.Msg) if err != nil { - t.Fatalf("RemoveBrevAuthorizedKeys: %v", err) + return nil, err } + return connect.NewResponse(resp), nil +} - if len(removed) != 2 { - t.Errorf("expected 2 removed keys, got %d: %v", len(removed), removed) - } +func (f *fakeNodeService) AddNode(_ context.Context, _ *connect.Request[nodev1.AddNodeRequest]) (*connect.Response[nodev1.AddNodeResponse], error) { + f.addNodeCalls++ + return connect.NewResponse(&nodev1.AddNodeResponse{}), nil +} - result := readAuthorizedKeys(t, u) - if strings.Contains(result, "#brev-portID:") { - t.Errorf("brev keys still present:\n%s", result) - } - if !strings.Contains(result, "ssh-rsa EXISTING user@host") { - t.Errorf("non-brev key was removed:\n%s", result) - } - if !strings.Contains(result, "ssh-ed25519 OTHERKEY admin@server") { - t.Errorf("non-brev key was removed:\n%s", result) +func startFakeServer(t *testing.T, svc *fakeNodeService) enableSSHDeps { + t.Helper() + _, handler := nodev1connect.NewExternalNodeServiceHandler(svc) + server := httptest.NewServer(handler) + t.Cleanup(server.Close) + return enableSSHDeps{ + nodeClients: mockNodeClientFactory{serverURL: server.URL}, } } -func Test_RemoveBrevAuthorizedKeys_NoopWhenFileDoesNotExist(t *testing.T) { - u := tempUser(t) +type enableSSHOrder struct{ entries []string } - removed, err := register.RemoveBrevAuthorizedKeys(u) - if err != nil { - t.Fatalf("expected no error for missing file, got: %v", err) - } - if len(removed) != 0 { - t.Errorf("expected no removed keys, got %v", removed) - } +func (o *enableSSHOrder) add(entry string) { o.entries = append(o.entries, entry) } + +type orderedPlatform struct{ order *enableSSHOrder } + +func (p orderedPlatform) IsCompatible() bool { + p.order.add("platform") + return true } -func Test_RemoveBrevAuthorizedKeys_NoopWhenNoBrevKeys(t *testing.T) { - u := tempUser(t) - sshDir := filepath.Join(u.HomeDir, ".ssh") - if err := os.MkdirAll(sshDir, 0o700); err != nil { - t.Fatal(err) - } +type orderedRegistrationStore struct { + order *enableSSHOrder + reg *register.DeviceRegistration + exists bool + err error +} - original := "ssh-rsa EXISTING user@host\nssh-ed25519 OTHER admin@server\n" - if err := os.WriteFile(filepath.Join(sshDir, "authorized_keys"), []byte(original), 0o600); err != nil { - t.Fatal(err) - } +func (s *orderedRegistrationStore) Save(*register.DeviceRegistration) error { + return errors.New("Save must not be called") +} - removed, err := register.RemoveBrevAuthorizedKeys(u) - if err != nil { - t.Fatalf("RemoveBrevAuthorizedKeys: %v", err) - } - if len(removed) != 0 { - t.Errorf("expected no removed keys, got %v", removed) - } +func (s *orderedRegistrationStore) Load() (*register.DeviceRegistration, error) { + return s.reg, s.err +} +func (s *orderedRegistrationStore) Delete() error { return errors.New("Delete must not be called") } +func (s *orderedRegistrationStore) Exists() (bool, error) { + s.order.add("registration") + return s.exists, s.err +} - result := readAuthorizedKeys(t, u) - if result != original { - t.Errorf("file was modified when it shouldn't have been.\nwant:\n%s\ngot:\n%s", original, result) - } +type orderedEnableSSHStore struct { + order *enableSSHOrder + user *entity.User } -// --- RemoveAuthorizedKey (specific key removal) --- +func (s orderedEnableSSHStore) GetCurrentUser() (*entity.User, error) { + s.order.add("auth") + return s.user, nil +} +func (orderedEnableSSHStore) GetAccessToken() (string, error) { return "token", nil } -func Test_RemoveAuthorizedKey_RemovesOnlyTargetKey(t *testing.T) { - u := tempUser(t) - sshDir := filepath.Join(u.HomeDir, ".ssh") - if err := os.MkdirAll(sshDir, 0o700); err != nil { - t.Fatal(err) - } +type orderedTunnel struct { + order *enableSSHOrder + err error +} - content := strings.Join([]string{ - "ssh-rsa KEEP1 user@host", - "ssh-rsa TARGET " + register.DevplaneAuthorizedKeysComment("p1", "u1"), - "ssh-rsa KEEP2 admin@server", - "", - }, "\n") - if err := os.WriteFile(filepath.Join(sshDir, "authorized_keys"), []byte(content), 0o600); err != nil { - t.Fatal(err) - } +func (t orderedTunnel) EnsureConnected(context.Context) error { + t.order.add("tunnel") + return t.err +} - if err := register.RemoveAuthorizedKey(u, "ssh-rsa TARGET"); err != nil { - t.Fatalf("RemoveAuthorizedKey: %v", err) - } +type reconnectingTunnel struct { + order *enableSSHOrder + connected *bool + reconnectAttempts int +} - result := readAuthorizedKeys(t, u) - if strings.Contains(result, "TARGET") { - t.Errorf("target key still present:\n%s", result) - } - if !strings.Contains(result, "ssh-rsa KEEP1 user@host") { - t.Errorf("unrelated key was removed:\n%s", result) - } - if !strings.Contains(result, "ssh-rsa KEEP2 admin@server") { - t.Errorf("unrelated key was removed:\n%s", result) +func (t *reconnectingTunnel) EnsureConnected(context.Context) error { + t.order.add("tunnel") + if !*t.connected { + t.reconnectAttempts++ + *t.connected = true } + return nil } -func Test_RemoveAuthorizedKey_NoopWhenKeyNotPresent(t *testing.T) { - u := tempUser(t) - sshDir := filepath.Join(u.HomeDir, ".ssh") - if err := os.MkdirAll(sshDir, 0o700); err != nil { - t.Fatal(err) - } +type orderedProvisioner struct { + order *enableSSHOrder + err error +} - original := "ssh-rsa EXISTING user@host\n" - if err := os.WriteFile(filepath.Join(sshDir, "authorized_keys"), []byte(original), 0o600); err != nil { - t.Fatal(err) - } +func (p orderedProvisioner) Provision( + context.Context, + *terminal.Terminal, + externalnode.TokenProvider, + *register.DeviceRegistration, + *entity.User, + *nodev1.ExternalNode, +) error { + p.order.add("provision") + return p.err +} - if err := register.RemoveAuthorizedKey(u, "ssh-rsa NOTHERE"); err != nil { - t.Fatalf("RemoveAuthorizedKey: %v", err) - } +type connectedTunnelProvisioner struct { + order *enableSSHOrder + tunnelConnected *bool + observedConnected bool +} - result := readAuthorizedKeys(t, u) - if !strings.Contains(result, "ssh-rsa EXISTING user@host") { - t.Errorf("existing key was removed:\n%s", result) - } +func (p *connectedTunnelProvisioner) Provision( + context.Context, + *terminal.Terminal, + externalnode.TokenProvider, + *register.DeviceRegistration, + *entity.User, + *nodev1.ExternalNode, +) error { + p.order.add("provision") + p.observedConnected = *p.tunnelConnected + if !p.observedConnected { + return errors.New("SSH provisioning started before the Brev tunnel connected") + } + return nil } -func Test_RemoveAuthorizedKey_NoopCases(t *testing.T) { - tests := []struct { - name string - key string - }{ - {"MissingFile", "ssh-rsa SOMEKEY"}, - {"EmptyKey", ""}, - {"WhitespaceKey", " "}, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - u := tempUser(t) - if err := register.RemoveAuthorizedKey(u, tt.key); err != nil { - t.Fatalf("expected no error, got: %v", err) - } - }) +func newEnableSSHTestDeps(order *enableSSHOrder, factory externalnode.NodeClientFactory, registrationStore register.RegistrationStore, tunnelErr error) enableSSHDeps { + return enableSSHDeps{ + platform: orderedPlatform{order: order}, + nodeClients: factory, + registrationStore: registrationStore, + tunnel: orderedTunnel{order: order, err: tunnelErr}, + provisioner: orderedProvisioner{order: order}, } } -func Test_RemoveAuthorizedKey_DoesNotRemoveOtherBrevKeys(t *testing.T) { - u := tempUser(t) - sshDir := filepath.Join(u.HomeDir, ".ssh") - if err := os.MkdirAll(sshDir, 0o700); err != nil { - t.Fatal(err) - } +func TestNewCmdEnableSSH_RejectsPositionalArguments(t *testing.T) { + cmd := NewCmdEnableSSH(terminal.New(), &mockEnableSSHStore{}) + require.Error(t, cmd.Args(cmd, []string{"unexpected"})) +} - content := strings.Join([]string{ - "ssh-rsa ALICE_KEY " + register.DevplaneAuthorizedKeysComment("p1", "u1"), - "ssh-rsa BOB_KEY " + register.DevplaneAuthorizedKeysComment("p2", "u2"), - "", - }, "\n") - if err := os.WriteFile(filepath.Join(sshDir, "authorized_keys"), []byte(content), 0o600); err != nil { - t.Fatal(err) - } +func TestRunEnableSSH_MissingRegistrationDirectsUserToJoin(t *testing.T) { + order := &enableSSHOrder{} + registrationStore := &orderedRegistrationStore{order: order, exists: false} + deps := newEnableSSHTestDeps(order, nil, registrationStore, nil) - // Remove only Alice's key — Bob's should stay. - if err := register.RemoveAuthorizedKey(u, "ssh-rsa ALICE_KEY"); err != nil { - t.Fatalf("RemoveAuthorizedKey: %v", err) - } + err := runEnableSSH(context.Background(), terminal.New(), orderedEnableSSHStore{order: order}, deps) - result := readAuthorizedKeys(t, u) - if strings.Contains(result, "ALICE_KEY") { - t.Errorf("Alice's key still present:\n%s", result) - } - if !strings.Contains(result, "ssh-rsa BOB_KEY") { - t.Errorf("Bob's key was removed:\n%s", result) - } + require.EqualError(t, err, `This machine has not joined a Brev network; run "brev join" first.`) + require.Equal(t, []string{"platform", "registration"}, order.entries) } -type mockNodeClientFactory struct{ serverURL string } +func TestRunEnableSSH_MissingBackendNodeDoesNotConnectOrProvision(t *testing.T) { + order := &enableSSHOrder{} + svc := &fakeNodeService{ + order: &order.entries, + getNodeFn: func(*nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) { + return &nodev1.GetNodeResponse{}, nil + }, + } + deps := startFakeServer(t, svc) + registrationStore := &orderedRegistrationStore{order: order, exists: true, reg: ®ister.DeviceRegistration{ExternalNodeID: "unode_123", OrgID: "org_456"}} + deps.platform = orderedPlatform{order: order} + deps.registrationStore = registrationStore + deps.tunnel = orderedTunnel{order: order} + deps.provisioner = orderedProvisioner{order: order} -func (m mockNodeClientFactory) NewNodeClient(provider externalnode.TokenProvider, _ string) nodev1connect.ExternalNodeServiceClient { - return register.NewNodeServiceClient(provider, m.serverURL) -} + err := runEnableSSH(context.Background(), terminal.New(), orderedEnableSSHStore{order: order, user: &entity.User{ID: "user_123"}}, deps) -type mockEnableSSHStore struct { - token string + require.ErrorContains(t, err, "registered node was not returned by Brev") + require.Equal(t, []string{"platform", "registration", "auth", "node"}, order.entries) } -func (m *mockEnableSSHStore) GetCurrentUser() (interface{}, error) { return nil, nil } -func (m *mockEnableSSHStore) GetAccessToken() (string, error) { return m.token, nil } +func TestRunEnableSSH_ConnectedTunnelProvisionsSSH(t *testing.T) { + order := &enableSSHOrder{} + svc := &fakeNodeService{ + order: &order.entries, + getNodeFn: func(*nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) { + return &nodev1.GetNodeResponse{ExternalNode: &nodev1.ExternalNode{ExternalNodeId: "unode_123"}}, nil + }, + } + deps := startFakeServer(t, svc) + registrationStore := &orderedRegistrationStore{order: order, exists: true, reg: ®ister.DeviceRegistration{ExternalNodeID: "unode_123", OrgID: "org_456", DisplayName: "joined-node"}} + deps.platform = orderedPlatform{order: order} + deps.registrationStore = registrationStore + deps.tunnel = orderedTunnel{order: order} + deps.provisioner = orderedProvisioner{order: order} -// fakeNodeService implements the server side of ExternalNodeService for testing. -type fakeNodeService struct { - nodev1connect.UnimplementedExternalNodeServiceHandler - getNodeFn func(*nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) + err := runEnableSSH(context.Background(), terminal.New(), orderedEnableSSHStore{order: order, user: &entity.User{ID: "user_123"}}, deps) + + require.NoError(t, err) + require.Equal(t, []string{"platform", "registration", "auth", "node", "tunnel", "provision"}, order.entries) } -func (f *fakeNodeService) GetNode(_ context.Context, req *connect.Request[nodev1.GetNodeRequest]) (*connect.Response[nodev1.GetNodeResponse], error) { - resp, err := f.getNodeFn(req.Msg) - if err != nil { - return nil, err +func TestRunEnableSSH_ReconnectsBeforeProvisioning(t *testing.T) { + order := &enableSSHOrder{} + tunnelConnected := false + svc := &fakeNodeService{ + order: &order.entries, + getNodeFn: func(*nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) { + return &nodev1.GetNodeResponse{ExternalNode: &nodev1.ExternalNode{ExternalNodeId: "unode_123"}}, nil + }, } - return connect.NewResponse(resp), nil -} + deps := startFakeServer(t, svc) + deps.platform = orderedPlatform{order: order} + deps.registrationStore = &orderedRegistrationStore{order: order, exists: true, reg: ®ister.DeviceRegistration{ExternalNodeID: "unode_123", OrgID: "org_456"}} + tunnel := &reconnectingTunnel{order: order, connected: &tunnelConnected} + deps.tunnel = tunnel + provisioner := &connectedTunnelProvisioner{order: order, tunnelConnected: &tunnelConnected} + deps.provisioner = provisioner -func startFakeServer(t *testing.T, svc *fakeNodeService) (enableSSHDeps, *httptest.Server) { - t.Helper() - _, handler := nodev1connect.NewExternalNodeServiceHandler(svc) - server := httptest.NewServer(handler) - t.Cleanup(server.Close) - return enableSSHDeps{ - nodeClients: mockNodeClientFactory{serverURL: server.URL}, - }, server + err := runEnableSSH(context.Background(), terminal.New(), orderedEnableSSHStore{order: order, user: &entity.User{ID: "user_123"}}, deps) + + require.NoError(t, err) + require.Equal(t, 1, tunnel.reconnectAttempts) + require.True(t, provisioner.observedConnected) + require.Equal(t, []string{"platform", "registration", "auth", "node", "tunnel", "provision"}, order.entries) } -func Test_fetchRegisteredNode(t *testing.T) { - svc := &fakeNodeService{ - getNodeFn: func(req *nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) { - if req.GetExternalNodeId() != "unode_abc" { - t.Fatalf("unexpected node id %q", req.GetExternalNodeId()) - } - return &nodev1.GetNodeResponse{ExternalNode: &nodev1.ExternalNode{ - ExternalNodeId: "unode_abc", - Ports: []*nodev1.Port{{PortId: "port_1", PortNumber: 11640, ServerPort: 22}}, - }}, nil +func TestRunEnableSSH_TunnelFailureDoesNotProvision(t *testing.T) { + tests := []struct { + name string + tunnelErr error + wantErrMsg string + }{ + { + name: "generic reconnect failure", + tunnelErr: errors.New("tunnel failed"), + wantErrMsg: "enable SSH requires a connected Brev tunnel", + }, + { + name: "connection remains unconfirmed", + tunnelErr: errors.New("Brev tunnel connection was not confirmed"), + wantErrMsg: "Brev tunnel connection was not confirmed", }, } - deps, _ := startFakeServer(t, svc) - store := &mockEnableSSHStore{token: "tok"} - reg := ®ister.DeviceRegistration{ExternalNodeID: "unode_abc", OrgID: "org_1"} - node, err := fetchRegisteredNode(context.Background(), deps, store, reg) - if err != nil { - t.Fatal(err) + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + order := &enableSSHOrder{} + svc := &fakeNodeService{ + order: &order.entries, + getNodeFn: func(*nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) { + return &nodev1.GetNodeResponse{ExternalNode: &nodev1.ExternalNode{ExternalNodeId: "unode_123"}}, nil + }, + } + deps := startFakeServer(t, svc) + deps.platform = orderedPlatform{order: order} + deps.registrationStore = &orderedRegistrationStore{order: order, exists: true, reg: ®ister.DeviceRegistration{ExternalNodeID: "unode_123", OrgID: "org_456"}} + deps.tunnel = orderedTunnel{order: order, err: tt.tunnelErr} + deps.provisioner = orderedProvisioner{order: order} + + err := runEnableSSH(context.Background(), terminal.New(), orderedEnableSSHStore{order: order, user: &entity.User{ID: "user_123"}}, deps) + + require.ErrorContains(t, err, tt.wantErrMsg) + require.NotContains(t, order.entries, "provision") + }) } - if len(node.GetPorts()) != 1 || node.GetPorts()[0].GetPortId() != "port_1" { - t.Fatalf("unexpected node: %+v", node) +} + +func TestRunEnableSSH_NeverAddsNode(t *testing.T) { + order := &enableSSHOrder{} + svc := &fakeNodeService{ + order: &order.entries, + getNodeFn: func(*nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) { + return &nodev1.GetNodeResponse{ExternalNode: &nodev1.ExternalNode{ExternalNodeId: "unode_123"}}, nil + }, } + deps := startFakeServer(t, svc) + deps.platform = orderedPlatform{order: order} + deps.registrationStore = &orderedRegistrationStore{order: order, exists: true, reg: ®ister.DeviceRegistration{ExternalNodeID: "unode_123", OrgID: "org_456"}} + deps.tunnel = orderedTunnel{order: order} + deps.provisioner = orderedProvisioner{order: order} + + err := runEnableSSH(context.Background(), terminal.New(), orderedEnableSSHStore{order: order, user: &entity.User{ID: "user_123"}}, deps) + + require.NoError(t, err) + require.Zero(t, svc.addNodeCalls) } diff --git a/pkg/cmd/register/device_registration_store.go b/pkg/cmd/register/device_registration_store.go index 315dfb99f..133443146 100644 --- a/pkg/cmd/register/device_registration_store.go +++ b/pkg/cmd/register/device_registration_store.go @@ -80,7 +80,7 @@ func (s *FileRegistrationStore) Load() (*DeviceRegistration, error) { if err != nil { return nil, breverrors.WrapAndTrace(err) } - return nil, breverrors.New("device registration not found, run 'brev register' first") + return nil, breverrors.New("device registration not found, run 'brev join' first") } var reg DeviceRegistration if err := files.ReadJSON(files.AppFs, path, ®); err != nil { diff --git a/pkg/cmd/register/device_registration_store_test.go b/pkg/cmd/register/device_registration_store_test.go index 39d7b1a21..2fb399e78 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" @@ -140,7 +141,10 @@ func Test_LoadRegistration_FailsWhenMissing(t *testing.T) { _, err := store.Load() if err == nil { - t.Error("expected error loading missing registration") + t.Fatal("expected error loading missing registration") + } + if !strings.Contains(err.Error(), "brev join") || strings.Contains(err.Error(), "brev register") { + t.Errorf("expected join recovery guidance, got: %v", err) } } diff --git a/pkg/cmd/register/node.go b/pkg/cmd/register/node.go new file mode 100644 index 000000000..df524186a --- /dev/null +++ b/pkg/cmd/register/node.go @@ -0,0 +1,37 @@ +package register + +import ( + "context" + "fmt" + + nodev1 "buf.build/gen/go/brevdev/devplane/protocolbuffers/go/devplaneapi/v1" + "connectrpc.com/connect" + + "github.com/brevdev/brev-cli/pkg/config" + "github.com/brevdev/brev-cli/pkg/externalnode" +) + +// FetchRegisteredNode retrieves the backend node represented by a local joined-device registration. +func FetchRegisteredNode( + ctx context.Context, + nodeClients externalnode.NodeClientFactory, + tokenProvider externalnode.TokenProvider, + reg *DeviceRegistration, +) (*nodev1.ExternalNode, error) { + client := nodeClients.NewNodeClient(tokenProvider, config.GlobalConfig.GetBrevPublicAPIURL()) + resp, err := client.GetNode(ctx, connect.NewRequest(&nodev1.GetNodeRequest{ + ExternalNodeId: reg.ExternalNodeID, + OrganizationId: reg.OrgID, + })) + if err != nil { + return nil, fmt.Errorf("error retrieving joined node: %w", err) + } + return registeredNodeFromResponse(resp) +} + +func registeredNodeFromResponse(resp *connect.Response[nodev1.GetNodeResponse]) (*nodev1.ExternalNode, error) { + if resp == nil || resp.Msg == nil || resp.Msg.GetExternalNode() == nil { + return nil, fmt.Errorf(`registered node was not returned by Brev; run "brev leave" and "brev join" to repair membership`) + } + return resp.Msg.GetExternalNode(), nil +} diff --git a/pkg/cmd/register/node_test.go b/pkg/cmd/register/node_test.go new file mode 100644 index 000000000..1dabf234b --- /dev/null +++ b/pkg/cmd/register/node_test.go @@ -0,0 +1,98 @@ +package register + +import ( + "context" + "errors" + "net/http/httptest" + "testing" + + nodev1connect "buf.build/gen/go/brevdev/devplane/connectrpc/go/devplaneapi/v1/devplaneapiv1connect" + nodev1 "buf.build/gen/go/brevdev/devplane/protocolbuffers/go/devplaneapi/v1" + "connectrpc.com/connect" + "github.com/stretchr/testify/require" + + "github.com/brevdev/brev-cli/pkg/externalnode" +) + +type registeredNodeTestFactory struct{ serverURL string } + +func (f registeredNodeTestFactory) NewNodeClient(provider externalnode.TokenProvider, _ string) nodev1connect.ExternalNodeServiceClient { + return NewNodeServiceClient(provider, f.serverURL) +} + +type registeredNodeTestTokenProvider struct{} + +func (registeredNodeTestTokenProvider) GetAccessToken() (string, error) { return "token", nil } + +type registeredNodeTestService struct { + nodev1connect.UnimplementedExternalNodeServiceHandler + getNode func(*nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) +} + +func (s registeredNodeTestService) GetNode(_ context.Context, req *connect.Request[nodev1.GetNodeRequest]) (*connect.Response[nodev1.GetNodeResponse], error) { + resp, err := s.getNode(req.Msg) + if err != nil { + return nil, err + } + return connect.NewResponse(resp), nil +} + +func startRegisteredNodeTestServer(t *testing.T, service registeredNodeTestService) registeredNodeTestFactory { + t.Helper() + _, handler := nodev1connect.NewExternalNodeServiceHandler(service) + server := httptest.NewServer(handler) + t.Cleanup(server.Close) + return registeredNodeTestFactory{serverURL: server.URL} +} + +func TestFetchRegisteredNode_Success(t *testing.T) { + factory := startRegisteredNodeTestServer(t, registeredNodeTestService{ + getNode: func(req *nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) { + require.Equal(t, "unode_123", req.GetExternalNodeId()) + require.Equal(t, "org_456", req.GetOrganizationId()) + return &nodev1.GetNodeResponse{ExternalNode: &nodev1.ExternalNode{ExternalNodeId: "unode_123"}}, nil + }, + }) + + node, err := FetchRegisteredNode(context.Background(), factory, registeredNodeTestTokenProvider{}, &DeviceRegistration{ + ExternalNodeID: "unode_123", + OrgID: "org_456", + }) + + require.NoError(t, err) + require.Equal(t, "unode_123", node.GetExternalNodeId()) +} + +func TestFetchRegisteredNode_RPCError(t *testing.T) { + factory := startRegisteredNodeTestServer(t, registeredNodeTestService{ + getNode: func(*nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) { + return nil, connect.NewError(connect.CodeInternal, errors.New("backend unavailable")) + }, + }) + + _, err := FetchRegisteredNode(context.Background(), factory, registeredNodeTestTokenProvider{}, &DeviceRegistration{ + ExternalNodeID: "unode_123", + OrgID: "org_456", + }) + + require.Error(t, err) + require.ErrorContains(t, err, "error retrieving joined node") +} + +func TestFetchRegisteredNode_NilNodeIsError(t *testing.T) { + tests := []struct { + name string + resp *connect.Response[nodev1.GetNodeResponse] + }{ + {name: "nil response"}, + {name: "nil message", resp: &connect.Response[nodev1.GetNodeResponse]{}}, + {name: "nil node", resp: connect.NewResponse(&nodev1.GetNodeResponse{})}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, err := registeredNodeFromResponse(tt.resp) + require.EqualError(t, err, `registered node was not returned by Brev; run "brev leave" and "brev join" to repair membership`) + }) + } +} diff --git a/pkg/cmd/register/providers.go b/pkg/cmd/register/providers.go index a3c27bd52..7d5754021 100644 --- a/pkg/cmd/register/providers.go +++ b/pkg/cmd/register/providers.go @@ -1,10 +1,13 @@ package register import ( + "context" "fmt" + "os" "os/exec" "runtime" "strings" + "time" nodev1connect "buf.build/gen/go/brevdev/devplane/connectrpc/go/devplaneapi/v1/devplaneapiv1connect" @@ -35,37 +38,130 @@ func (TerminalPrompter) Select(label string, items []string) string { }) } -// Netbird handles NetBird installation and uninstallation. -type Netbird struct{} +func (TerminalPrompter) Input(content terminal.PromptContent) string { + return terminal.PromptGetInput(content) +} + +const ( + defaultNetBirdConnectTimeout = 30 * time.Second + defaultNetBirdPollInterval = 500 * time.Millisecond +) + +type netBirdCommandRunner interface { + Output(context.Context, string, ...string) ([]byte, error) + Run(context.Context, string, ...string) error +} + +type execNetBirdCommandRunner struct{} + +func (execNetBirdCommandRunner) Output(ctx context.Context, name string, args ...string) ([]byte, error) { + return exec.CommandContext(ctx, name, args...).Output() //nolint:wrapcheck // EnsureConnected adds operation context. +} + +func (execNetBirdCommandRunner) Run(ctx context.Context, name string, args ...string) error { + cmd := exec.CommandContext(ctx, name, args...) + cmd.Stdin = os.Stdin + cmd.Stdout = os.Stdout + cmd.Stderr = os.Stderr + return cmd.Run() //nolint:wrapcheck // EnsureConnected adds operation context. +} + +// Netbird handles NetBird installation and connectivity. +type Netbird struct { + runner netBirdCommandRunner + connectTimeout time.Duration + pollInterval time.Duration +} func (Netbird) Install() error { return InstallNetbird() } func (Netbird) Uninstall() error { return UninstallNetbird() } -// EnsureRunning checks if the netbird systemd service is active and attempts -// to start it if it is not. It also checks the netbird peer connection status -// and runs "netbird up" if the peer is disconnected. -func (Netbird) EnsureRunning() error { - out, err := exec.Command("systemctl", "is-active", "netbird").Output() //nolint:gosec // fixed service name +func (n Netbird) commandRunner() netBirdCommandRunner { + if n.runner != nil { + return n.runner + } + return execNetBirdCommandRunner{} +} + +func (n Netbird) connectionTimeout() time.Duration { + if n.connectTimeout > 0 { + return n.connectTimeout + } + return defaultNetBirdConnectTimeout +} + +func (n Netbird) connectionPollInterval() time.Duration { + if n.pollInterval > 0 { + return n.pollInterval + } + return defaultNetBirdPollInterval +} + +// EnsureConnected ensures the local service is active and confirms that its +// management connection is established before returning. +func (n Netbird) EnsureConnected(ctx context.Context) error { + runner := n.commandRunner() + out, err := runner.Output(ctx, "systemctl", "is-active", "netbird") if err != nil || strings.TrimSpace(string(out)) != "active" { - if startErr := exec.Command("sudo", "systemctl", "start", "netbird").Run(); startErr != nil { //nolint:gosec // fixed service name + if startErr := runner.Run(ctx, "sudo", "systemctl", "start", "netbird"); startErr != nil { return fmt.Errorf("failed to start Brev tunnel service: %w", startErr) } } - statusOut, err := exec.Command("netbird", "status").Output() //nolint:gosec // fixed command - if err != nil { - // Service is running, just can't confirm peer status. + statusOut, statusErr := runner.Output(ctx, "netbird", "status") + if statusErr == nil && netbirdManagementConnected(string(statusOut)) { return nil } - if netbirdManagementConnected(string(statusOut)) { + if upErr := runner.Run(ctx, "sudo", "netbird", "up"); upErr != nil { + return fmt.Errorf("failed to reconnect Brev tunnel: %w", upErr) + } + + confirmationCtx, cancel := context.WithTimeout(ctx, n.connectionTimeout()) + defer cancel() + lastStatusErr := statusErr + checkStatus := func() bool { + statusOut, err := runner.Output(confirmationCtx, "netbird", "status") + if err != nil { + lastStatusErr = err + return false + } + return netbirdManagementConnected(string(statusOut)) + } + + if checkStatus() { return nil } - if upErr := exec.Command("sudo", "netbird", "up").Run(); upErr != nil { //nolint:gosec // fixed command - return fmt.Errorf("failed to reconnect Brev tunnel: %w", upErr) + ticker := time.NewTicker(n.connectionPollInterval()) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return fmt.Errorf("wait for Brev tunnel connection: %w", ctx.Err()) + case <-confirmationCtx.Done(): + if lastStatusErr != nil { + return fmt.Errorf("Brev tunnel connection was not confirmed: %w", lastStatusErr) + } + return fmt.Errorf("Brev tunnel connection was not confirmed: %w", confirmationCtx.Err()) + case <-ticker.C: + if checkStatus() { + return nil + } + } + } +} + +// netbirdManagementConnected parses "netbird status" output and returns true +// when the Management line reports "Connected". +func netbirdManagementConnected(statusOutput string) bool { + for _, line := range strings.Split(statusOutput, "\n") { + line = strings.TrimSpace(line) + if strings.HasPrefix(line, "Management:") { + return strings.TrimSpace(strings.TrimPrefix(line, "Management:")) == "Connected" + } } - return nil + return false } // ShellSetupRunner runs setup scripts via shell. diff --git a/pkg/cmd/register/providers_test.go b/pkg/cmd/register/providers_test.go new file mode 100644 index 000000000..a5f92a47a --- /dev/null +++ b/pkg/cmd/register/providers_test.go @@ -0,0 +1,177 @@ +package register + +import ( + "context" + "errors" + "reflect" + "strings" + "testing" + "time" +) + +type netBirdCall struct { + name string + args []string +} + +type netBirdResult struct { + output []byte + err error +} + +type fakeNetBirdCommandRunner struct { + results []netBirdResult + fallback netBirdResult + calls []netBirdCall +} + +func (f *fakeNetBirdCommandRunner) Output(_ context.Context, name string, args ...string) ([]byte, error) { + f.calls = append(f.calls, netBirdCall{name: name, args: append([]string(nil), args...)}) + if len(f.results) == 0 { + return append([]byte(nil), f.fallback.output...), f.fallback.err + } + result := f.results[0] + f.results = f.results[1:] + return append([]byte(nil), result.output...), result.err +} + +func (f *fakeNetBirdCommandRunner) Run(ctx context.Context, name string, args ...string) error { + _, err := f.Output(ctx, name, args...) + return err +} + +func connectedNetBirdStatus() []byte { + return []byte("Management: Connected\n") +} + +func disconnectedNetBirdStatus() []byte { + return []byte("Management: Disconnected\n") +} + +func newTestNetbird(runner *fakeNetBirdCommandRunner) Netbird { + return Netbird{ + runner: runner, + connectTimeout: 10 * time.Millisecond, + pollInterval: time.Millisecond, + } +} + +func TestNetbirdEnsureConnected_AlreadyConnectedDoesNotReconnect(t *testing.T) { + runner := &fakeNetBirdCommandRunner{results: []netBirdResult{ + {output: []byte("active\n")}, + {output: connectedNetBirdStatus()}, + }} + + err := newTestNetbird(runner).EnsureConnected(context.Background()) + if err != nil { + t.Fatalf("EnsureConnected() error = %v", err) + } + + wantCalls := []netBirdCall{ + {name: "systemctl", args: []string{"is-active", "netbird"}}, + {name: "netbird", args: []string{"status"}}, + } + if !reflect.DeepEqual(runner.calls, wantCalls) { + t.Fatalf("commands = %#v, want %#v", runner.calls, wantCalls) + } +} + +func TestNetbirdEnsureConnected_StartsInactiveService(t *testing.T) { + runner := &fakeNetBirdCommandRunner{results: []netBirdResult{ + {output: []byte("inactive\n")}, + {}, + {output: connectedNetBirdStatus()}, + }} + + err := newTestNetbird(runner).EnsureConnected(context.Background()) + if err != nil { + t.Fatalf("EnsureConnected() error = %v", err) + } + + wantCalls := []netBirdCall{ + {name: "systemctl", args: []string{"is-active", "netbird"}}, + {name: "sudo", args: []string{"systemctl", "start", "netbird"}}, + {name: "netbird", args: []string{"status"}}, + } + if !reflect.DeepEqual(runner.calls, wantCalls) { + t.Fatalf("commands = %#v, want %#v", runner.calls, wantCalls) + } +} + +func TestNetbirdEnsureConnected_ReconnectsAndWaitsForConfirmation(t *testing.T) { + runner := &fakeNetBirdCommandRunner{results: []netBirdResult{ + {output: []byte("active\n")}, + {output: disconnectedNetBirdStatus()}, + {}, + {output: disconnectedNetBirdStatus()}, + {output: connectedNetBirdStatus()}, + }} + + err := newTestNetbird(runner).EnsureConnected(context.Background()) + if err != nil { + t.Fatalf("EnsureConnected() error = %v", err) + } + + wantCalls := []netBirdCall{ + {name: "systemctl", args: []string{"is-active", "netbird"}}, + {name: "netbird", args: []string{"status"}}, + {name: "sudo", args: []string{"netbird", "up"}}, + {name: "netbird", args: []string{"status"}}, + {name: "netbird", args: []string{"status"}}, + } + if !reflect.DeepEqual(runner.calls, wantCalls) { + t.Fatalf("commands = %#v, want %#v", runner.calls, wantCalls) + } +} + +func TestNetbirdEnsureConnected_ReconnectFailure(t *testing.T) { + runner := &fakeNetBirdCommandRunner{results: []netBirdResult{ + {output: []byte("active\n")}, + {output: disconnectedNetBirdStatus()}, + {err: errors.New("up failed")}, + }} + + err := newTestNetbird(runner).EnsureConnected(context.Background()) + if err == nil || !strings.Contains(err.Error(), "failed to reconnect Brev tunnel") { + t.Fatalf("EnsureConnected() error = %v, want reconnect failure context", err) + } +} + +func TestNetbirdEnsureConnected_StatusNeverConfirmsConnection(t *testing.T) { + runner := &fakeNetBirdCommandRunner{ + results: []netBirdResult{ + {output: []byte("active\n")}, + {output: disconnectedNetBirdStatus()}, + {}, + }, + fallback: netBirdResult{output: disconnectedNetBirdStatus()}, + } + + err := newTestNetbird(runner).EnsureConnected(context.Background()) + if err == nil || !strings.Contains(err.Error(), "Brev tunnel connection was not confirmed") { + t.Fatalf("EnsureConnected() error = %v, want confirmation timeout", err) + } + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("EnsureConnected() error = %v, want context deadline exceeded", err) + } +} + +func TestNetbirdEnsureConnected_StatusErrorsAreNotSuccess(t *testing.T) { + statusErr := errors.New("netbird status unavailable") + runner := &fakeNetBirdCommandRunner{ + results: []netBirdResult{ + {output: []byte("active\n")}, + {err: statusErr}, + {}, + }, + fallback: netBirdResult{err: statusErr}, + } + + err := newTestNetbird(runner).EnsureConnected(context.Background()) + if err == nil { + t.Fatal("EnsureConnected() error = nil, want confirmation timeout") + } + if !strings.Contains(err.Error(), "Brev tunnel connection was not confirmed") || !strings.Contains(err.Error(), statusErr.Error()) { + t.Fatalf("EnsureConnected() error = %v, want timeout containing latest status failure", err) + } +} diff --git a/pkg/cmd/register/register.go b/pkg/cmd/register/register.go index 2ad88b435..424c5248a 100644 --- a/pkg/cmd/register/register.go +++ b/pkg/cmd/register/register.go @@ -1,11 +1,10 @@ -// Package register provides the brev register command for device registration +// Package register provides the brev join command and device registration storage. package register import ( "context" "errors" "fmt" - "os/user" "strings" "time" @@ -34,14 +33,16 @@ type RegisterStore interface { GetAccessToken() (string, error) } +// NetBirdConnector confirms local NetBird management connectivity. +type NetBirdConnector interface { + EnsureConnected(context.Context) error +} + // NetBirdManager installs, uninstalls, and monitors the NetBird network agent. type NetBirdManager interface { + NetBirdConnector Install() error Uninstall() error - // EnsureRunning checks whether the NetBird service is active and - // connected, starting or reconnecting it if needed. Returns nil when - // the tunnel is healthy. - EnsureRunning() error } // SetupRunner runs a setup script on the local machine. @@ -49,12 +50,17 @@ type SetupRunner interface { RunSetup(script string) error } -// registerDeps bundles the side-effecting dependencies of runRegister so they +type joinPrompter interface { + terminal.Confirmer + terminal.Selector + Input(terminal.PromptContent) string +} + +// joinDeps bundles the side-effecting dependencies of runJoin so they // can be replaced in tests. -type registerDeps struct { +type joinDeps struct { platform externalnode.PlatformChecker - prompter terminal.Confirmer - selector terminal.Selector + prompter joinPrompter gater sudo.Gater netbird NetBirdManager setupRunner SetupRunner @@ -63,12 +69,11 @@ type registerDeps struct { registrationStore RegistrationStore } -func defaultRegisterDeps() registerDeps { +func defaultJoinDeps() joinDeps { p := TerminalPrompter{} - return registerDeps{ + return joinDeps{ platform: LinuxPlatform{}, prompter: p, - selector: p, gater: sudo.Default, netbird: Netbird{}, setupRunner: ShellSetupRunner{}, @@ -79,23 +84,34 @@ func defaultRegisterDeps() registerDeps { } var ( - registerLong = `Register your device with NVIDIA Brev + joinLong = `Join this device to a Brev network -This command sets up network connectivity and registers this machine with Brev. +This command sets up network connectivity and joins this machine to Brev. 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 join' with no flags and follow prompts for device name and organization. + • 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 + joinExample = ` # Interactive (prompts for device name, organization, and confirmations) + brev join # 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` + brev join --name my-node --org my-org` ) +func NewCmdJoin(t *terminal.Terminal, store RegisterStore) *cobra.Command { + return newCmdJoin(t, store, defaultJoinDeps) +} + +// NewCmdRegister is retained for source compatibility. It returns the +// canonical join command with register as its deprecated alias. +// +// Deprecated: use NewCmdJoin. func NewCmdRegister(t *terminal.Terminal, store RegisterStore) *cobra.Command { + return NewCmdJoin(t, store) +} + +func newCmdJoin(t *terminal.Terminal, store RegisterStore, depsFactory func() joinDeps) *cobra.Command { var orgFlag string var nameFlag string var sshPort int @@ -103,50 +119,56 @@ func NewCmdRegister(t *terminal.Terminal, store RegisterStore) *cobra.Command { cmd := &cobra.Command{ Annotations: map[string]string{"configuration": ""}, - Use: "register", + Use: "join", + Aliases: []string{"register"}, DisableFlagsInUseLine: true, - Short: "Register this device with Brev", - Long: registerLong, - Example: registerExample, + Short: "Join this device to a Brev network", + Long: joinLong, + Example: joinExample, Args: cobra.NoArgs, RunE: func(cmd *cobra.Command, args []string) error { - interactive := nameFlag == "" && orgFlag == "" && sshPort == 0 - opts := registerOpts{ - interactive: interactive, + if cmd.CalledAs() == "register" { + fmt.Fprintln(cmd.ErrOrStderr(), `Warning: "brev register" is deprecated; use "brev join" instead.`) + fmt.Fprintln(cmd.ErrOrStderr(), `This command no longer enables SSH; run "brev enable-ssh" separately.`) + } + if cmd.Flags().Changed("ssh-port") { + return fmt.Errorf("--ssh-port is no longer supported by brev join or brev register; run brev join, then run brev enable-ssh on the joined machine") + } + opts := joinOpts{ + interactive: nameFlag == "" && orgFlag == "", name: nameFlag, orgName: orgFlag, - sshPort: int32(sshPort), skipConfirm: approveFlag, } - return runRegister(cmd.Context(), t, store, opts, defaultRegisterDeps()) + return runJoin(cmd.Context(), t, store, opts, depsFactory()) }, } cmd.Flags().StringVarP(&orgFlag, "org", "o", "", "organization name (required when using non-interactive mode)") 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().IntVarP(&sshPort, "ssh-port", "p", 0, "deprecated") + _ = cmd.Flags().MarkHidden("ssh-port") cmd.Flags().BoolVar(&approveFlag, "approve", false, "skip all confirmation prompts (assume yes)") return cmd } -// registerOpts carries mode and inputs: when interactive, name/orgName/sshPort are from prompts; otherwise from flags. -type registerOpts struct { +// joinOpts carries mode and inputs: when interactive, name and orgName are prompted; otherwise they come from flags. +type joinOpts 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 +// runJoin runs a single membership setup flow; the only difference by mode is whether we prompt or use opts. +func runJoin(ctx context.Context, t *terminal.Terminal, s RegisterStore, opts joinOpts, deps joinDeps) error { //nolint:gocognit,gocyclo,funlen // ok // Basic validation if !deps.platform.IsCompatible() { - return breverrors.New("brev register is only supported on Linux") + return breverrors.New("brev join is only supported on Linux") } // Always gate on sudo; skip confirmation prompt when non-interactive or --approve. - if err := deps.gater.Gate(t, deps.prompter, "Device registration", !opts.interactive || opts.skipConfirm); err != nil { + if err := deps.gater.Gate(t, deps.prompter, "Device join", !opts.interactive || opts.skipConfirm); err != nil { return fmt.Errorf("sudo issue: %w", err) } if !opts.interactive { @@ -156,7 +178,7 @@ func runRegister(ctx context.Context, t *terminal.Terminal, s RegisterStore, opt } // Run through the login flow - brevUser, err := s.GetCurrentUser() + _, err := s.GetCurrentUser() if err != nil { return breverrors.WrapAndTrace(err) } @@ -174,7 +196,7 @@ func runRegister(ctx context.Context, t *terminal.Terminal, s RegisterStore, opt var name string if opts.interactive { t.Vprint("") - name = terminal.PromptGetInput(terminal.PromptContent{ + name = deps.prompter.Input(terminal.PromptContent{ Label: "Device name", ErrorMsg: "name is required", AllowEmpty: false, @@ -201,7 +223,7 @@ func runRegister(ctx context.Context, t *terminal.Terminal, s RegisterStore, opt t.Vprint("") t.Vprint(t.White("══════════════════════════════════════════════════")) - t.Vprint(t.White(" Registering your device with Brev")) + t.Vprint(t.White(" Joining your device to Brev")) t.Vprint(t.White("══════════════════════════════════════════════════")) t.Vprint("") if opts.interactive && !opts.skipConfirm { @@ -214,60 +236,34 @@ func runRegister(ctx context.Context, t *terminal.Terminal, s RegisterStore, opt t.Vprint(t.Yellow(" This will:")) t.Vprint(" 1. Download and install Brev tunnel") t.Vprint(" 2. Collect hardware profile") - t.Vprint(" 3. Register this machine with Brev") - t.Vprint(" 4. Store registration data") + t.Vprint(" 3. Join this machine to Brev") + t.Vprint(" 4. Store join data") t.Vprint(" 5. Connect device to Brev") t.Vprint("") if opts.interactive { - if !opts.skipConfirm && !deps.prompter.ConfirmYesNo("Proceed with registration?") { - t.Vprint("Registration canceled.") + if !opts.skipConfirm && !deps.prompter.ConfirmYesNo("Proceed with join?") { + t.Vprint("Join canceled.") return nil } } - // Perform the registration steps - reg, err := runRegisterSteps(ctx, t, s, name, org, deps) - if err != nil { + if err := runJoinSteps(ctx, t, s, name, org, deps); 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))) - } - } - + t.Vprint("") + t.Vprint("SSH access was not enabled. To enable it for your user, run: brev enable-ssh") 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) { +// runJoinSteps performs netbird install, hardware profile, AddNode, save registration, and runSetup. +func runJoinSteps(ctx context.Context, t *terminal.Terminal, s RegisterStore, name string, org *entity.Organization, deps joinDeps) 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 +271,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("") @@ -283,7 +279,7 @@ func runRegisterSteps(ctx context.Context, t *terminal.Terminal, s RegisterStore t.Vprint(FormatHardwareProfile(hwProfile)) t.Vprint("") - t.Vprint(t.Yellow("[Step 3/5] Registering device with Brev...")) + t.Vprint(t.Yellow("[Step 3/5] Joining device to Brev...")) deviceID := uuid.New().String() client := deps.nodeClients.NewNodeClient(s, config.GlobalConfig.GetBrevPublicAPIURL()) addResp, err := client.AddNode(ctx, connect.NewRequest(&nodev1.AddNodeRequest{ @@ -297,9 +293,9 @@ func runRegisterSteps(ctx context.Context, t *terminal.Terminal, s RegisterStore // 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()) + return errors.New(connectErr.Message()) } - return nil, fmt.Errorf("failed to register node: %w", err) + return fmt.Errorf("failed to join node: %w", err) } node := addResp.Msg.GetExternalNode() @@ -316,24 +312,24 @@ func runRegisterSteps(ctx context.Context, t *terminal.Terminal, s RegisterStore 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 joined but failed to save locally: %w", err) } t.Vprint("") t.Vprint(t.Yellow("[Step 5/5] Connecting device to Brev...")) runSetup(node, t, deps) - t.Vprintf("%s Node registered.\n", t.Green(" ✓")) - t.Vprintf("%s Registration complete.\n", t.Green(" ✓")) - return reg, nil + t.Vprintf("%s Node joined.\n", t.Green(" ✓")) + t.Vprintf("%s Join complete.\n", t.Green(" ✓")) + return nil } -func resolveOrgInteractive(t *terminal.Terminal, s RegisterStore, deps registerDeps) (*entity.Organization, error) { +func resolveOrgInteractive(t *terminal.Terminal, s RegisterStore, deps joinDeps) (*entity.Organization, error) { list, err := s.ListOrganizations() if err != nil { return nil, breverrors.WrapAndTrace(err) } - org, err := helpers.SelectOrganizationInteractive(t, list, deps.selector) + org, err := helpers.SelectOrganizationInteractive(t, list, deps.prompter) if err != nil { return nil, breverrors.WrapAndTrace(err) } @@ -352,7 +348,7 @@ func resolveOrg(s RegisterStore, orgName string) (*entity.Organization, error) { // 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 { +func checkExistingRegistration(ctx context.Context, t *terminal.Terminal, s RegisterStore, deps joinDeps) 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) @@ -377,39 +373,25 @@ func checkExistingRegistration(ctx context.Context, t *terminal.Terminal, s Regi ci := node.GetConnectivityInfo() if ci != nil && ci.GetStatus() == nodev1.NetworkMemberStatus_NETWORK_MEMBER_STATUS_CONNECTED { t.Vprint(t.Green(" Node is connected.")) - t.Vprint("") - t.Vprint(" Run 'brev deregister' first if you want to re-register.") - return nil + } else { + t.Vprintf(" Node status: %s\n", externalnode.FriendlyNetworkStatus(ci.GetStatus())) } - t.Vprintf(" Node status: %s\n", externalnode.FriendlyNetworkStatus(ci.GetStatus())) } - // Check local netbird service and start it if down. + // Confirm local NetBird connectivity even when the backend is connected. t.Vprint(" Checking local Brev tunnel...") - if err := deps.netbird.EnsureRunning(); err != nil { + if err := deps.netbird.EnsureConnected(ctx); err != nil { t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: %v", err))) } else { - t.Vprint(t.Green(" Brev tunnel is running.")) + t.Vprint(t.Green(" Brev tunnel is connected.")) } t.Vprint("") - t.Vprint(" Run 'brev deregister' first if you want to re-register.") + t.Vprint(" Run 'brev leave' first if you want to rejoin.") return nil } -// netbirdManagementConnected parses "netbird status" output and returns true -// when the Management line reports "Connected". -func netbirdManagementConnected(statusOutput string) bool { - for _, line := range strings.Split(statusOutput, "\n") { - line = strings.TrimSpace(line) - if strings.HasPrefix(line, "Management:") { - return strings.TrimSpace(strings.TrimPrefix(line, "Management:")) == "Connected" - } - } - return false -} - -func runSetup(node *nodev1.ExternalNode, t *terminal.Terminal, deps registerDeps) { +func runSetup(node *nodev1.ExternalNode, t *terminal.Terminal, deps joinDeps) { ci := node.GetConnectivityInfo() if ci == nil || ci.GetRegistrationCommand() == "" { t.Vprintf(" %s\n", t.Yellow("Warning: Brev tunnel setup failed, please try again.")) @@ -423,59 +405,3 @@ 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 - } - - t.Vprint("") - t.Vprint(t.White("══════════════════════════════════════════════════")) - t.Vprint(t.White(" Enabling SSH access on this device")) - 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.Vprint("") - - 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 -} diff --git a/pkg/cmd/register/register_test.go b/pkg/cmd/register/register_test.go index d98b1a92a..44e463b1f 100644 --- a/pkg/cmd/register/register_test.go +++ b/pkg/cmd/register/register_test.go @@ -1,9 +1,12 @@ package register import ( + "bytes" "context" "fmt" + "io" "net/http/httptest" + "os" "strings" "testing" @@ -15,8 +18,205 @@ import ( "github.com/brevdev/brev-cli/pkg/externalnode" "github.com/brevdev/brev-cli/pkg/sudo" "github.com/brevdev/brev-cli/pkg/terminal" + "github.com/spf13/cobra" + "github.com/stretchr/testify/require" ) +const legacySSHPortMigrationError = "--ssh-port is no longer supported by brev join or brev register; run brev join, then run brev enable-ssh on the joined machine" + +type panicRegisterStore struct{} + +func (panicRegisterStore) GetCurrentUser() (*entity.User, error) { panic("GetCurrentUser called") } +func (panicRegisterStore) GetActiveOrganizationOrDefault() (*entity.Organization, error) { + panic("GetActiveOrganizationOrDefault called") +} + +func (panicRegisterStore) GetOrganizationsByName(string) ([]entity.Organization, error) { + panic("GetOrganizationsByName called") +} + +func (panicRegisterStore) ListOrganizations() ([]entity.Organization, error) { + panic("ListOrganizations called") +} +func (panicRegisterStore) GetAccessToken() (string, error) { panic("GetAccessToken called") } + +func TestNewCmdJoin_CommandSurface(t *testing.T) { + cmd := NewCmdJoin(terminal.New(), panicRegisterStore{}) + root := &cobra.Command{Use: "brev"} + root.AddCommand(cmd) + + resolved, _, err := root.Find([]string{"register"}) + require.NoError(t, err) + require.Equal(t, "join", cmd.Name()) + require.Equal(t, []string{"register"}, cmd.Aliases) + require.Same(t, cmd, resolved) + require.Error(t, cmd.Args(cmd, []string{"unexpected"})) + require.True(t, cmd.Flags().Lookup("ssh-port").Hidden) +} + +func TestNewCmdRegister_DeprecatedSourceCompatibility(t *testing.T) { + cmd := NewCmdRegister(terminal.New(), panicRegisterStore{}) + + require.Equal(t, "join", cmd.Name()) + require.Equal(t, []string{"register"}, cmd.Aliases) +} + +func TestNewCmdJoin_RegisterAliasWarnsOnExecution(t *testing.T) { + cmd := NewCmdJoin(terminal.New(), panicRegisterStore{}) + root := &cobra.Command{Use: "brev", SilenceUsage: true} + root.AddCommand(cmd) + var stderr bytes.Buffer + root.SetErr(&stderr) + root.SetArgs([]string{"register", "--ssh-port", "22"}) + + err := root.Execute() + + require.EqualError(t, err, legacySSHPortMigrationError) + require.Contains(t, stderr.String(), "Warning: \"brev register\" is deprecated; use \"brev join\" instead.\nThis command no longer enables SSH; run \"brev enable-ssh\" separately.\n") +} + +func TestNewCmdJoin_HelpDoesNotWarn(t *testing.T) { + cmd := NewCmdJoin(terminal.New(), panicRegisterStore{}) + root := &cobra.Command{Use: "brev", SilenceUsage: true} + root.AddCommand(cmd) + var stderr bytes.Buffer + root.SetErr(&stderr) + root.SetArgs([]string{"register", "--help"}) + + require.NoError(t, root.Execute()) + require.Empty(t, stderr.String()) +} + +func TestNewCmdJoin_LegacySSHPortFailsBeforeSideEffects(t *testing.T) { + tests := [][]string{ + {"join", "--ssh-port", "0"}, + {"join", "--ssh-port", "22"}, + {"join", "-p", "0"}, + {"join", "-p", "22"}, + {"register", "--ssh-port", "0"}, + {"register", "--ssh-port", "22"}, + {"register", "-p", "0"}, + {"register", "-p", "22"}, + } + + for _, args := range tests { + t.Run(strings.Join(args, " "), func(t *testing.T) { + depsConstructed := 0 + cmd := newCmdJoin(terminal.New(), panicRegisterStore{}, func() joinDeps { + depsConstructed++ + return joinDeps{} + }) + root := &cobra.Command{Use: "brev", SilenceUsage: true} + root.AddCommand(cmd) + root.SetErr(&bytes.Buffer{}) + root.SetArgs(args) + + require.EqualError(t, root.Execute(), legacySSHPortMigrationError) + // All platform, sudo, authentication, NetBird, RPC, persistence, + // setup, and hardware work is contained in the dependency factory. + require.Zero(t, depsConstructed) + }) + } +} + +type recordingJoinPrompter struct { + prompts []joinPrompt +} + +type joinPrompt struct { + kind string + label string +} + +func (p *recordingJoinPrompter) ConfirmYesNo(label string) bool { + p.prompts = append(p.prompts, joinPrompt{kind: "confirm", label: label}) + return true +} + +func (p *recordingJoinPrompter) Select(label string, items []string) string { + p.prompts = append(p.prompts, joinPrompt{kind: "select", label: label}) + return items[0] +} + +func (p *recordingJoinPrompter) Input(content terminal.PromptContent) string { + p.prompts = append(p.prompts, joinPrompt{kind: "input", label: content.Label}) + return "interactive-node" +} + +func TestRunJoin_InteractivePromptsOnlyForMembership(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: req.GetOrganizationId(), Name: req.GetName(), DeviceId: req.GetDeviceId(), + }}, nil + }} + deps, server := testJoinDeps(t, svc, regStore) + defer server.Close() + prompter := &recordingJoinPrompter{} + deps.prompter = prompter + + require.NoError(t, runJoin(context.Background(), terminal.New(), store, joinOpts{interactive: true}, deps)) + require.Equal(t, []joinPrompt{ + {kind: "input", label: "Device name"}, + {kind: "select", label: "Select organization"}, + {kind: "confirm", label: "Proceed with join?"}, + }, prompter.prompts) +} + +func TestRunJoin_DoesNotOpenPortOrGrantSSH(t *testing.T) { + regStore := &mockRegistrationStore{} + store := &mockRegisterStore{ + user: &entity.User{ID: "user_1"}, + org: &entity.Organization{ID: "org_123", Name: "TestOrg"}, + token: "tok", + } + openCalls, grantCalls := 0, 0 + svc := &fakeNodeService{ + addNodeFn: func(req *nodev1.AddNodeRequest) (*nodev1.AddNodeResponse, error) { + return &nodev1.AddNodeResponse{ExternalNode: &nodev1.ExternalNode{ + ExternalNodeId: "unode_abc", OrganizationId: req.GetOrganizationId(), Name: req.GetName(), DeviceId: req.GetDeviceId(), + }}, nil + }, + openPortFn: func(*nodev1.OpenPortRequest) (*nodev1.OpenPortResponse, error) { + openCalls++ + return &nodev1.OpenPortResponse{}, nil + }, + grantNodeSSHAccessFn: func(*nodev1.GrantNodeSSHAccessRequest) (*nodev1.GrantNodeSSHAccessResponse, error) { + grantCalls++ + return &nodev1.GrantNodeSSHAccessResponse{}, nil + }, + } + deps, server := testJoinDeps(t, svc, regStore) + defer server.Close() + + stdout := captureStdout(t) + require.NoError(t, runJoin(context.Background(), terminal.New(), store, joinOpts{name: "my-node", orgName: "TestOrg"}, deps)) + require.Equal(t, 0, openCalls) + require.Equal(t, 0, grantCalls) + require.Contains(t, stdout(), "brev enable-ssh") +} + +func captureStdout(t *testing.T) func() string { + t.Helper() + previous := os.Stdout + reader, writer, err := os.Pipe() + require.NoError(t, err) + os.Stdout = writer + return func() string { + require.NoError(t, writer.Close()) + os.Stdout = previous + output, err := io.ReadAll(reader) + require.NoError(t, err) + require.NoError(t, reader.Close()) + return string(output) + } +} + // mockRegisterStore satisfies RegisterStore for orchestration tests. type mockRegisterStore struct { user *entity.User @@ -88,7 +288,7 @@ func (m *mockRegistrationStore) Exists() (bool, error) { return m.reg != nil, nil } -// mock types for registerDeps interfaces +// mock types for joinDeps interfaces type mockPlatform struct{ compatible bool } @@ -96,30 +296,32 @@ func (m mockPlatform) IsCompatible() bool { return m.compatible } type mockConfirmer struct{ confirm bool } -func (m mockConfirmer) ConfirmYesNo(_ string) bool { return m.confirm } - -// mockSelector implements terminal.Selector by returning the first item (for tests that need org selection). -type mockSelector struct{ choice string } - -func (m mockSelector) Select(_ string, items []string) string { - if m.choice != "" { - for _, s := range items { - if s == m.choice { - return s - } - } - } - if len(items) > 0 { - return items[0] +func (m mockConfirmer) ConfirmYesNo(_ string) bool { return m.confirm } +func (m mockConfirmer) Input(_ terminal.PromptContent) string { return "" } +func (m mockConfirmer) Select(_ string, items []string) string { + if len(items) == 0 { + return "" } - return "" + return items[0] } type mockNetBirdManager struct{ err error } -func (m mockNetBirdManager) Install() error { return m.err } -func (m mockNetBirdManager) Uninstall() error { return m.err } -func (m mockNetBirdManager) EnsureRunning() error { return m.err } +func (m mockNetBirdManager) Install() error { return m.err } +func (m mockNetBirdManager) Uninstall() error { return m.err } +func (m mockNetBirdManager) EnsureConnected(context.Context) error { return m.err } + +type reconcilingNetBirdManager struct { + called bool + err error +} + +func (m *reconcilingNetBirdManager) Install() error { return m.err } +func (m *reconcilingNetBirdManager) Uninstall() error { return m.err } +func (m *reconcilingNetBirdManager) EnsureConnected(context.Context) error { + m.called = true + return m.err +} type mockSetupRunner struct { called bool @@ -161,18 +363,17 @@ func testHardwareProfile() *HardwareProfile { } } -// testRegisterDeps returns deps with all side effects stubbed out, and a fake +// testJoinDeps returns deps with all side effects stubbed out, and a fake // ConnectRPC server backed by the provided fakeNodeService. -func testRegisterDeps(t *testing.T, svc *fakeNodeService, regStore RegistrationStore) (registerDeps, *httptest.Server) { +func testJoinDeps(t *testing.T, svc *fakeNodeService, regStore RegistrationStore) (joinDeps, *httptest.Server) { t.Helper() _, handler := nodev1connect.NewExternalNodeServiceHandler(svc) server := httptest.NewServer(handler) - return registerDeps{ + return joinDeps{ platform: mockPlatform{compatible: true}, prompter: mockConfirmer{confirm: true}, - selector: mockSelector{}, gater: sudo.CachedGater{}, netbird: mockNetBirdManager{}, setupRunner: &mockSetupRunner{}, @@ -184,7 +385,7 @@ func testRegisterDeps(t *testing.T, svc *fakeNodeService, regStore RegistrationS }, server } -func Test_runRegister_HappyPath(t *testing.T) { +func Test_runJoin_HappyPath(t *testing.T) { regStore := &mockRegistrationStore{} store := &mockRegisterStore{ @@ -218,19 +419,16 @@ func Test_runRegister_HappyPath(t *testing.T) { setupRunner := &mockSetupRunner{} - deps, server := testRegisterDeps(t, svc, regStore) + deps, server := testJoinDeps(t, svc, regStore) defer server.Close() deps.setupRunner = setupRunner - 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) + opts := joinOpts{interactive: false, name: "my-spark", orgName: "TestOrg"} + err := runJoin(context.Background(), term, store, opts, deps) if err != nil { - t.Fatalf("runRegister failed: %v", err) + t.Fatalf("runJoin failed: %v", err) } // Verify registration was persisted @@ -239,7 +437,7 @@ func Test_runRegister_HappyPath(t *testing.T) { t.Fatalf("Exists error: %v", err) } if !exists { - t.Fatal("expected registration to exist after successful register") + t.Fatal("expected registration to exist after successful join") } reg, err := regStore.Load() @@ -262,14 +460,14 @@ func Test_runRegister_HappyPath(t *testing.T) { } } -// gaterFromFunc adapts a function to sudo.Gater; used only by Test_runRegister_UserCancels. +// gaterFromFunc adapts a function to sudo.Gater; used only by Test_runJoin_UserCancels. type gaterFromFunc func(*terminal.Terminal, terminal.Confirmer, string, bool) error func (f gaterFromFunc) Gate(t *terminal.Terminal, c terminal.Confirmer, reason string, assumeYes bool) error { return f(t, c, reason, assumeYes) } -func Test_runRegister_UserCancels(t *testing.T) { +func Test_runJoin_UserCancels(t *testing.T) { // User cancel happens in interactive mode (sudo or confirm). Flag-driven has no prompts. regStore := &mockRegistrationStore{} store := &mockRegisterStore{ @@ -278,7 +476,7 @@ func Test_runRegister_UserCancels(t *testing.T) { token: "tok", } svc := &fakeNodeService{} - deps, server := testRegisterDeps(t, svc, regStore) + deps, server := testJoinDeps(t, svc, regStore) defer server.Close() deps.prompter = mockConfirmer{confirm: false} @@ -295,8 +493,8 @@ func Test_runRegister_UserCancels(t *testing.T) { }) term := terminal.New() - opts := registerOpts{interactive: true, name: "", orgName: "", sshPort: 0} - err := runRegister(context.Background(), term, store, opts, deps) + opts := joinOpts{interactive: true, name: "", orgName: ""} + err := runJoin(context.Background(), term, store, opts, deps) if err == nil { t.Fatal("expected error when user declines sudo gate") } @@ -310,7 +508,7 @@ func Test_runRegister_UserCancels(t *testing.T) { } } -func Test_runRegister_AlreadyRegistered(t *testing.T) { +func Test_runJoin_AlreadyRegistered(t *testing.T) { tests := []struct { name string getNodeFn func(*nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) @@ -377,14 +575,14 @@ func Test_runRegister_AlreadyRegistered(t *testing.T) { } svc := &fakeNodeService{getNodeFn: tt.getNodeFn} - deps, server := testRegisterDeps(t, svc, regStore) + deps, server := testJoinDeps(t, svc, regStore) defer server.Close() 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} - err := runRegister(context.Background(), term, store, opts, deps) + opts := joinOpts{interactive: false, name: "Existing", orgName: "TestOrg"} + err := runJoin(context.Background(), term, store, opts, deps) if err != nil { t.Fatalf("expected nil error, got: %v", err) } @@ -397,7 +595,39 @@ func Test_runRegister_AlreadyRegistered(t *testing.T) { } } -func Test_runRegister_NoOrganization(t *testing.T) { +func TestCheckExistingRegistration_ReconcilesLocalTunnel(t *testing.T) { + regStore := &mockRegistrationStore{ + reg: &DeviceRegistration{ + ExternalNodeID: "unode_existing", + DisplayName: "Existing", + OrgID: "org_123", + }, + } + store := &mockRegisterStore{token: "tok"} + 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 + }} + deps, server := testJoinDeps(t, svc, regStore) + defer server.Close() + tunnel := &reconcilingNetBirdManager{} + deps.netbird = tunnel + + if err := checkExistingRegistration(context.Background(), terminal.New(), store, deps); err != nil { + t.Fatalf("checkExistingRegistration() error = %v", err) + } + if !tunnel.called { + t.Fatal("checkExistingRegistration() did not reconcile the local Brev tunnel") + } +} + +func Test_runJoin_NoOrganization(t *testing.T) { regStore := &mockRegistrationStore{} store := &mockRegisterStore{ @@ -408,18 +638,18 @@ func Test_runRegister_NoOrganization(t *testing.T) { } svc := &fakeNodeService{} - deps, server := testRegisterDeps(t, svc, regStore) + deps, server := testJoinDeps(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) + opts := joinOpts{interactive: false, name: "my-spark", orgName: "TestOrg"} + err := runJoin(context.Background(), term, store, opts, deps) if err == nil { t.Fatal("expected error when no org exists") } } -func Test_runRegister_WithOrgFlag(t *testing.T) { +func Test_runJoin_WithOrgFlag(t *testing.T) { regStore := &mockRegistrationStore{} store := &mockRegisterStore{ @@ -447,18 +677,15 @@ func Test_runRegister_WithOrgFlag(t *testing.T) { } setupRunner := &mockSetupRunner{} - deps, server := testRegisterDeps(t, svc, regStore) + deps, server := testJoinDeps(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) + opts := joinOpts{interactive: false, name: "my-spark", orgName: "SpecificOrg"} + err := runJoin(context.Background(), term, store, opts, deps) if err != nil { - t.Fatalf("runRegister with --org failed: %v", err) + t.Fatalf("runJoin with --org failed: %v", err) } if capturedOrgID != "org_456" { @@ -474,7 +701,7 @@ func Test_runRegister_WithOrgFlag(t *testing.T) { } } -func Test_runRegister_WithOrgFlag_NotFound(t *testing.T) { +func Test_runJoin_WithOrgFlag_NotFound(t *testing.T) { regStore := &mockRegistrationStore{} store := &mockRegisterStore{ @@ -485,12 +712,12 @@ func Test_runRegister_WithOrgFlag_NotFound(t *testing.T) { } svc := &fakeNodeService{} - deps, server := testRegisterDeps(t, svc, regStore) + deps, server := testJoinDeps(t, svc, regStore) defer server.Close() term := terminal.New() - opts := registerOpts{interactive: false, name: "my-spark", orgName: "NonexistentOrg", sshPort: 22} - err := runRegister(context.Background(), term, store, opts, deps) + opts := joinOpts{interactive: false, name: "my-spark", orgName: "NonexistentOrg"} + err := runJoin(context.Background(), term, store, opts, deps) if err == nil { t.Fatal("expected error when org not found") } @@ -499,7 +726,7 @@ func Test_runRegister_WithOrgFlag_NotFound(t *testing.T) { } } -func Test_runRegister_AddNodeFails(t *testing.T) { +func Test_runJoin_AddNodeFails(t *testing.T) { regStore := &mockRegistrationStore{} store := &mockRegisterStore{ @@ -515,12 +742,12 @@ func Test_runRegister_AddNodeFails(t *testing.T) { }, } - deps, server := testRegisterDeps(t, svc, regStore) + deps, server := testJoinDeps(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) + opts := joinOpts{interactive: false, name: "my-spark", orgName: "TestOrg"} + err := runJoin(context.Background(), term, store, opts, deps) if err == nil { t.Fatal("expected error when AddNode fails") } @@ -535,7 +762,7 @@ func Test_runRegister_AddNodeFails(t *testing.T) { } } -func Test_runRegister_NoSetupCommand(t *testing.T) { +func Test_runJoin_NoSetupCommand(t *testing.T) { regStore := &mockRegistrationStore{} store := &mockRegisterStore{ @@ -561,19 +788,16 @@ func Test_runRegister_NoSetupCommand(t *testing.T) { setupRunner := &mockSetupRunner{} - deps, server := testRegisterDeps(t, svc, regStore) + deps, server := testJoinDeps(t, svc, regStore) defer server.Close() deps.setupRunner = setupRunner - 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) + opts := joinOpts{interactive: false, name: "my-spark", orgName: "TestOrg"} + err := runJoin(context.Background(), term, store, opts, deps) if err != nil { - t.Fatalf("runRegister failed: %v", err) + t.Fatalf("runJoin failed: %v", err) } if setupRunner.called { @@ -662,110 +886,7 @@ 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", - } - - 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) - } - 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) - } -} - -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", - } - - 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) - }, - } - - 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 should not fail the overall flow when SSH grant fails: %v", err) - } - - if grantCalls != 1 { - t.Errorf("expected GrantNodeSSHAccess to be called once (no retry on permanent error), got %d", grantCalls) - } -} - -func Test_runRegister_NameValidation(t *testing.T) { +func Test_runJoin_NameValidation(t *testing.T) { tests := []struct { name string input string @@ -808,16 +929,13 @@ func Test_runRegister_NameValidation(t *testing.T) { }, } - deps, server := testRegisterDeps(t, svc, regStore) + deps, server := testJoinDeps(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) + opts := joinOpts{interactive: false, name: tt.input, orgName: "TestOrg"} + err = runJoin(context.Background(), term, store, opts, deps) if tt.wantErr { if err == nil { t.Fatal("expected error, got nil") @@ -832,7 +950,7 @@ func Test_runRegister_NameValidation(t *testing.T) { } } -func Test_runRegister_PlatformIncompatible(t *testing.T) { +func Test_runJoin_PlatformIncompatible(t *testing.T) { regStore := &mockRegistrationStore{} store := &mockRegisterStore{ @@ -842,14 +960,14 @@ func Test_runRegister_PlatformIncompatible(t *testing.T) { } svc := &fakeNodeService{} - deps, server := testRegisterDeps(t, svc, regStore) + deps, server := testJoinDeps(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) + opts := joinOpts{interactive: false, name: "my-spark", orgName: "TestOrg"} + err := runJoin(context.Background(), term, store, opts, deps) if err == nil { t.Fatal("expected error when platform is incompatible") } @@ -858,7 +976,7 @@ func Test_runRegister_PlatformIncompatible(t *testing.T) { } } -func Test_runRegister_HardwareProfilerFailure(t *testing.T) { +func Test_runJoin_HardwareProfilerFailure(t *testing.T) { regStore := &mockRegistrationStore{} store := &mockRegisterStore{ @@ -868,14 +986,14 @@ func Test_runRegister_HardwareProfilerFailure(t *testing.T) { } svc := &fakeNodeService{} - deps, server := testRegisterDeps(t, svc, regStore) + deps, server := testJoinDeps(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) + opts := joinOpts{interactive: false, name: "my-spark", orgName: "TestOrg"} + err := runJoin(context.Background(), term, store, opts, deps) if err == nil { t.Fatal("expected error when hardware profiler fails") } @@ -884,7 +1002,7 @@ func Test_runRegister_HardwareProfilerFailure(t *testing.T) { } } -func Test_runRegister_NetBirdInstallFailure(t *testing.T) { +func Test_runJoin_NetBirdInstallFailure(t *testing.T) { regStore := &mockRegistrationStore{} store := &mockRegisterStore{ @@ -894,14 +1012,14 @@ func Test_runRegister_NetBirdInstallFailure(t *testing.T) { } svc := &fakeNodeService{} - deps, server := testRegisterDeps(t, svc, regStore) + deps, server := testJoinDeps(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) + opts := joinOpts{interactive: false, name: "my-spark", orgName: "TestOrg"} + err := runJoin(context.Background(), term, store, opts, deps) if err == nil { t.Fatal("expected error when NetBird install fails") } @@ -910,7 +1028,7 @@ func Test_runRegister_NetBirdInstallFailure(t *testing.T) { } } -func Test_runRegister_NoNameNotRegistered(t *testing.T) { +func Test_runJoin_NoNameNotRegistered(t *testing.T) { // In flag-driven mode, missing --name and --org must error (no prompts). regStore := &mockRegistrationStore{} @@ -921,12 +1039,12 @@ func Test_runRegister_NoNameNotRegistered(t *testing.T) { } svc := &fakeNodeService{} - deps, server := testRegisterDeps(t, svc, regStore) + deps, server := testJoinDeps(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) + opts := joinOpts{interactive: false, name: "", orgName: ""} + err := runJoin(context.Background(), term, store, opts, deps) if err == nil { t.Fatal("expected error when no name/org in non-interactive mode") } @@ -935,7 +1053,7 @@ func Test_runRegister_NoNameNotRegistered(t *testing.T) { } } -func Test_runRegister_NoNameAlreadyRegistered(t *testing.T) { +func Test_runJoin_NoNameAlreadyRegistered(t *testing.T) { regStore := &mockRegistrationStore{ reg: &DeviceRegistration{ ExternalNodeID: "unode_existing", @@ -963,12 +1081,12 @@ func Test_runRegister_NoNameAlreadyRegistered(t *testing.T) { }, } - deps, server := testRegisterDeps(t, svc, regStore) + deps, server := testJoinDeps(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) + opts := joinOpts{interactive: false, name: "Existing", orgName: "TestOrg"} + err := runJoin(context.Background(), term, store, opts, deps) if err != nil { t.Fatalf("expected nil error when already registered with no name, got: %v", err) } @@ -979,153 +1097,3 @@ func Test_runRegister_NoNameAlreadyRegistered(t *testing.T) { t.Error("expected registration to still exist") } } - -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) - }) - } -} diff --git a/pkg/cmd/register/sshkeys.go b/pkg/cmd/register/sshkeys.go index 1188766dc..5f75c5150 100644 --- a/pkg/cmd/register/sshkeys.go +++ b/pkg/cmd/register/sshkeys.go @@ -134,98 +134,12 @@ const ( backoffPrintRound = 500 * time.Millisecond ) -// BrevKeyPrefixLegacy marks keys written by older CLI versions (# brev-cli). -const BrevKeyPrefixLegacy = "# brev-cli" - -// BrevKeyPrefix is an alias for BrevKeyPrefixLegacy (tests and migration). -const BrevKeyPrefix = BrevKeyPrefixLegacy - -const ( - brevPortIDField = "brev-portID:" - brevUserIDField = "brev-userID:" -) - // DevplaneAuthorizedKeysComment is the suffix on Brev-managed authorized_keys lines. // The CLI writes this before GrantNodeSSHAccess so devplane need not modify the file. func DevplaneAuthorizedKeysComment(portID, userID string) string { return fmt.Sprintf("#brev-portID:%s,brev-userID:%s", portID, userID) } -// BrevAuthorizedKey represents a single Brev-managed key found in authorized_keys. -type BrevAuthorizedKey struct { - Line string // full line from authorized_keys - KeyContent string // key type + material (and optional ssh comment), without brev suffix - PortID string // from devplane #brev-portID:... - UserID string // from devplane brev-userID:... or legacy user_id= -} - -func isBrevManagedAuthorizedKeysLine(line string) bool { - return strings.Contains(line, BrevKeyPrefixLegacy) || strings.Contains(line, "#brev-portID:") -} - -func parseBrevAuthorizedKeyLine(trimmed string) BrevAuthorizedKey { - bk := BrevAuthorizedKey{Line: trimmed} - - if idx := strings.Index(trimmed, "#brev-portID:"); idx >= 0 { - bk.KeyContent = strings.TrimSpace(trimmed[:idx]) - tag := trimmed[idx+1:] - for _, part := range strings.Split(tag, ",") { - part = strings.TrimSpace(part) - switch { - case strings.HasPrefix(part, brevPortIDField): - bk.PortID = strings.TrimPrefix(part, brevPortIDField) - case strings.HasPrefix(part, brevUserIDField): - bk.UserID = strings.TrimPrefix(part, brevUserIDField) - } - } - return bk - } - - if idx := strings.Index(trimmed, " "+BrevKeyPrefixLegacy); idx >= 0 { - bk.KeyContent = strings.TrimSpace(trimmed[:idx]) - tag := trimmed[idx+1:] - if uidIdx := strings.Index(tag, "user_id="); uidIdx >= 0 { - rest := tag[uidIdx+len("user_id="):] - if spIdx := strings.Index(rest, " "); spIdx >= 0 { - bk.UserID = rest[:spIdx] - } else { - bk.UserID = rest - } - } - return bk - } - - bk.KeyContent = trimmed - return bk -} - -// ListBrevAuthorizedKeys reads ~/.ssh/authorized_keys and returns Brev-managed lines. -func ListBrevAuthorizedKeys(u *user.User) ([]BrevAuthorizedKey, error) { - authKeysPath := filepath.Join(u.HomeDir, ".ssh", "authorized_keys") - - data, err := os.ReadFile(authKeysPath) // #nosec G304 - if err != nil { - if os.IsNotExist(err) { - return nil, nil - } - return nil, fmt.Errorf("reading authorized_keys: %w", err) - } - - var keys []BrevAuthorizedKey - for _, line := range strings.Split(string(data), "\n") { - if !isBrevManagedAuthorizedKeysLine(line) { - continue - } - trimmed := strings.TrimSpace(line) - if trimmed == "" { - continue - } - keys = append(keys, parseBrevAuthorizedKeyLine(trimmed)) - } - - return keys, nil -} - // RemoveAuthorizedKeyLine removes an exact line from authorized_keys. func RemoveAuthorizedKeyLine(u *user.User, line string) error { line = strings.TrimSpace(line) @@ -479,68 +393,3 @@ func InstallAuthorizedKey(u *user.User, pubKey, portID, brevUserID string) (bool return true, nil } - -// RemoveAuthorizedKey removes a specific public key from the user's -// ~/.ssh/authorized_keys. It matches the key content regardless of whether -// the brev-cli comment tag is present. -func RemoveAuthorizedKey(u *user.User, pubKey string) error { - pubKey = strings.TrimSpace(pubKey) - if pubKey == "" { - return nil - } - - authKeysPath := filepath.Join(u.HomeDir, ".ssh", "authorized_keys") - - existing, err := os.ReadFile(authKeysPath) // #nosec G304 - if err != nil { - if os.IsNotExist(err) { - return nil - } - return fmt.Errorf("reading authorized_keys: %w", err) - } - - var kept []string - for _, line := range strings.Split(string(existing), "\n") { - if strings.Contains(line, pubKey) { - continue - } - kept = append(kept, line) - } - - result := strings.Join(kept, "\n") - if err := os.WriteFile(authKeysPath, []byte(result), 0o600); err != nil { - return fmt.Errorf("writing authorized_keys: %w", err) - } - return nil -} - -// RemoveBrevAuthorizedKeys removes all Brev-managed SSH keys from authorized_keys. -func RemoveBrevAuthorizedKeys(u *user.User) ([]string, error) { - authKeysPath := filepath.Join(u.HomeDir, ".ssh", "authorized_keys") - - existing, err := os.ReadFile(authKeysPath) // #nosec G304 - if err != nil { - if os.IsNotExist(err) { - return nil, nil - } - return nil, fmt.Errorf("reading authorized_keys: %w", err) - } - - var kept []string - var removed []string - for _, line := range strings.Split(string(existing), "\n") { - if isBrevManagedAuthorizedKeysLine(line) { - if trimmed := strings.TrimSpace(line); trimmed != "" { - removed = append(removed, trimmed) - } - continue - } - kept = append(kept, line) - } - - result := strings.Join(kept, "\n") - if err := os.WriteFile(authKeysPath, []byte(result), 0o600); err != nil { - return nil, fmt.Errorf("writing authorized_keys: %w", err) - } - return removed, nil -} diff --git a/pkg/cmd/register/sshkeys_test.go b/pkg/cmd/register/sshkeys_test.go index acecb74ab..52d0ef203 100644 --- a/pkg/cmd/register/sshkeys_test.go +++ b/pkg/cmd/register/sshkeys_test.go @@ -43,87 +43,6 @@ func TestDevplaneAuthorizedKeysComment(t *testing.T) { } } -func TestListBrevAuthorizedKeys_ParsesDevplaneFormat(t *testing.T) { - u := tempUser(t) - seedKeys(t, u, strings.Join([]string{ - "ssh-rsa EXISTING user@host", - "ssh-ed25519 AAAA_ALICE user@a.com " + DevplaneAuthorizedKeysComment("port_1", "user_1"), - "ssh-rsa AAAA_BOB " + DevplaneAuthorizedKeysComment("port_2", "user_2"), - "", - }, "\n")) - - keys, err := ListBrevAuthorizedKeys(u) - if err != nil { - t.Fatalf("ListBrevAuthorizedKeys: %v", err) - } - if len(keys) != 2 { - t.Fatalf("expected 2 keys, got %d", len(keys)) - } - if keys[0].PortID != "port_1" || keys[0].UserID != "user_1" { - t.Errorf("key[0]: port=%q user=%q", keys[0].PortID, keys[0].UserID) - } - if keys[1].PortID != "port_2" || keys[1].UserID != "user_2" { - t.Errorf("key[1]: port=%q user=%q", keys[1].PortID, keys[1].UserID) - } -} - -func TestListBrevAuthorizedKeys_ParsesLegacyFormat(t *testing.T) { - u := tempUser(t) - seedKeys(t, u, "ssh-ed25519 AAAA_OLD # brev-cli user_id=uid_42\n") - - keys, err := ListBrevAuthorizedKeys(u) - if err != nil { - t.Fatalf("ListBrevAuthorizedKeys: %v", err) - } - if len(keys) != 1 { - t.Fatalf("expected 1 key, got %d", len(keys)) - } - if keys[0].UserID != "uid_42" { - t.Errorf("expected user_id uid_42, got %q", keys[0].UserID) - } -} - -func TestListBrevAuthorizedKeys_MixedFormats(t *testing.T) { - u := tempUser(t) - seedKeys(t, u, strings.Join([]string{ - "ssh-rsa AAAA_LEGACY # brev-cli", - "ssh-rsa NONBREV user@host", - "ssh-ed25519 AAAA_NEW " + DevplaneAuthorizedKeysComment("p1", "uid_42"), - "", - }, "\n")) - - keys, err := ListBrevAuthorizedKeys(u) - if err != nil { - t.Fatalf("ListBrevAuthorizedKeys: %v", err) - } - if len(keys) != 2 { - t.Fatalf("expected 2 brev keys, got %d", len(keys)) - } -} - -func TestListBrevAuthorizedKeys_NoFile(t *testing.T) { - u := tempUser(t) - keys, err := ListBrevAuthorizedKeys(u) - if err != nil { - t.Fatalf("expected no error for missing file, got: %v", err) - } - if len(keys) != 0 { - t.Errorf("expected 0 keys, got %d", len(keys)) - } -} - -func TestListBrevAuthorizedKeys_NoBrevKeys(t *testing.T) { - u := tempUser(t) - seedKeys(t, u, "ssh-rsa NONBREV user@host\n") - keys, err := ListBrevAuthorizedKeys(u) - if err != nil { - t.Fatalf("ListBrevAuthorizedKeys: %v", err) - } - if len(keys) != 0 { - t.Errorf("expected 0 brev keys, got %d", len(keys)) - } -} - func TestRemoveAuthorizedKeyLine_RemovesExactLine(t *testing.T) { u := tempUser(t) line := "ssh-ed25519 REMOVE " + DevplaneAuthorizedKeysComment("p1", "user_1") @@ -141,43 +60,6 @@ func TestRemoveAuthorizedKeyLine_RemovesExactLine(t *testing.T) { } } -func TestRemoveBrevAuthorizedKeys_DevplaneLines(t *testing.T) { - u := tempUser(t) - seedKeys(t, u, strings.Join([]string{ - "ssh-rsa KEEP user@host", - "ssh-rsa BREV1 " + DevplaneAuthorizedKeysComment("p1", "u1"), - "ssh-rsa BREV2 " + DevplaneAuthorizedKeysComment("p2", "u2"), - "", - }, "\n")) - - removed, err := RemoveBrevAuthorizedKeys(u) - if err != nil { - t.Fatalf("RemoveBrevAuthorizedKeys: %v", err) - } - if len(removed) != 2 { - t.Fatalf("expected 2 removed, got %d", len(removed)) - } - result := readKeys(t, u) - if strings.Contains(result, "#brev-portID:") { - t.Errorf("brev keys remain:\n%s", result) - } - if !strings.Contains(result, "KEEP") { - t.Error("non-brev key was removed") - } -} - -func TestRemoveBrevAuthorizedKeys_LegacyLines(t *testing.T) { - u := tempUser(t) - seedKeys(t, u, "ssh-rsa BREVKEY "+BrevKeyPrefixLegacy+"\n") - removed, err := RemoveBrevAuthorizedKeys(u) - if err != nil { - t.Fatalf("RemoveBrevAuthorizedKeys: %v", err) - } - if len(removed) != 1 { - t.Fatalf("expected 1 removed, got %d", len(removed)) - } -} - func TestInstallAuthorizedKey_AppendsDevplaneComment(t *testing.T) { u := tempUser(t) pub := "ssh-rsa AAAA testkey user@example.com" @@ -243,18 +125,6 @@ func TestInstallAuthorizedKey_secondPortAppendsNewLine(t *testing.T) { } } -func TestRemoveAuthorizedKey_ByPublicKeyMaterial(t *testing.T) { - u := tempUser(t) - pub := "ssh-rsa AAAA testkey" - seedKeys(t, u, pub+" user@host "+DevplaneAuthorizedKeysComment("p1", "u1")+"\n") - if err := RemoveAuthorizedKey(u, pub); err != nil { - t.Fatal(err) - } - if strings.Contains(readKeys(t, u), "AAAA") { - t.Fatal("key material should be removed") - } -} - // --- PromptSSHPort --- func TestPromptSSHPort(t *testing.T) { diff --git a/pkg/sudo/sudo.go b/pkg/sudo/sudo.go index d75c18ef1..9518e0963 100644 --- a/pkg/sudo/sudo.go +++ b/pkg/sudo/sudo.go @@ -29,10 +29,14 @@ type Gater interface { var Default Gater = &systemGater{} // systemGater implements Gater using the real sudo check and optional password prompt. -type systemGater struct{} +type systemGater struct { + checkStatus func() Status + runCommand func(*exec.Cmd) error + stdin *os.File +} func (g *systemGater) Gate(t *terminal.Terminal, confirmer terminal.Confirmer, reason string, assumeYes bool) error { - status := check() + status := g.status() if status == StatusRoot { return nil } @@ -51,15 +55,17 @@ func (g *systemGater) Gate(t *terminal.Terminal, confirmer terminal.Confirmer, r } if status == StatusUncached { - if exec.Command("sudo", "-n", "-v").Run() != nil { //nolint:gosec // intentional sudo -n -v - if isTTY(os.Stdin) { - cmd := exec.Command("sudo", "-v") //nolint:gosec // intentional sudo -v - cmd.Stdin = os.Stdin - cmd.Stdout = os.Stdout - cmd.Stderr = os.Stderr - if err := cmd.Run(); err != nil { - return fmt.Errorf("sudo authentication failed: %w", err) - } + if err := g.run(exec.Command("sudo", "-n", "-v")); err != nil { //nolint:gosec // intentional sudo -n -v + stdin := g.input() + if !isTTY(stdin) { + return fmt.Errorf("sudo authentication unavailable without an interactive terminal: %w", err) + } + cmd := exec.Command("sudo", "-v") //nolint:gosec // intentional sudo -v + cmd.Stdin = stdin + cmd.Stdout = os.Stdout + cmd.Stderr = os.Stderr + if err := g.run(cmd); err != nil { + return fmt.Errorf("sudo authentication failed: %w", err) } } } @@ -67,6 +73,27 @@ func (g *systemGater) Gate(t *terminal.Terminal, confirmer terminal.Confirmer, r return nil } +func (g *systemGater) status() Status { + if g.checkStatus != nil { + return g.checkStatus() + } + return check() +} + +func (g *systemGater) run(cmd *exec.Cmd) error { + if g.runCommand != nil { + return g.runCommand(cmd) + } + return cmd.Run() //nolint:wrapcheck // Gate adds branch-specific sudo authentication context. +} + +func (g *systemGater) input() *os.File { + if g.stdin != nil { + return g.stdin + } + return os.Stdin +} + // check returns the current sudo status. func check() Status { if os.Getuid() == 0 { diff --git a/pkg/sudo/sudo_test.go b/pkg/sudo/sudo_test.go new file mode 100644 index 000000000..0ed02b515 --- /dev/null +++ b/pkg/sudo/sudo_test.go @@ -0,0 +1,38 @@ +package sudo + +import ( + "errors" + "os" + "os/exec" + "testing" + + "github.com/brevdev/brev-cli/pkg/terminal" + "github.com/stretchr/testify/require" +) + +type sudoTestConfirmer struct{} + +func (sudoTestConfirmer) ConfirmYesNo(string) bool { return true } + +func TestSystemGater_UncachedNonInteractiveSudoFailureIsReturned(t *testing.T) { + stdin, err := os.CreateTemp(t.TempDir(), "stdin") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, stdin.Close()) }) + + probeErr := errors.New("sudo credentials unavailable") + runCalls := 0 + gater := &systemGater{ + checkStatus: func() Status { return StatusUncached }, + runCommand: func(cmd *exec.Cmd) error { + runCalls++ + require.Equal(t, []string{"sudo", "-n", "-v"}, cmd.Args) + return probeErr + }, + stdin: stdin, + } + + err = gater.Gate(terminal.New(), sudoTestConfirmer{}, "Leave Brev network", true) + require.ErrorIs(t, err, probeErr) + require.ErrorContains(t, err, "sudo authentication unavailable without an interactive terminal") + require.Equal(t, 1, runCalls) +}