diff --git a/go.mod b/go.mod index 5855b1103..442b4f383 100644 --- a/go.mod +++ b/go.mod @@ -165,3 +165,5 @@ require ( sigs.k8s.io/structured-merge-diff/v4 v4.4.1 // indirect sigs.k8s.io/yaml v1.4.0 // indirect ) + +replace buf.build/gen/go/brevdev/devplane/protocolbuffers/go => /tmp/buf-proto-patch diff --git a/pkg/cmd/allowssh/allowssh.go b/pkg/cmd/allowssh/allowssh.go new file mode 100644 index 000000000..60e997b6f --- /dev/null +++ b/pkg/cmd/allowssh/allowssh.go @@ -0,0 +1,213 @@ +// Package allowssh implements brev allow-ssh. +package allowssh + +import ( + "context" + "fmt" + "os" + "os/exec" + "os/user" + "path/filepath" + "strings" + + nodev1 "buf.build/gen/go/brevdev/devplane/protocolbuffers/go/devplaneapi/v1" + "connectrpc.com/connect" + + "github.com/brevdev/brev-cli/pkg/cmd/register" + "github.com/brevdev/brev-cli/pkg/config" + "github.com/brevdev/brev-cli/pkg/entity" + "github.com/brevdev/brev-cli/pkg/externalnode" + "github.com/brevdev/brev-cli/pkg/sshcert" + "github.com/brevdev/brev-cli/pkg/terminal" + + "github.com/spf13/cobra" +) + +type AllowSSHStore interface { + GetCurrentUser() (*entity.User, error) + GetAccessToken() (string, error) +} + +type allowSSHDeps struct { + platform externalnode.PlatformChecker + nodeClients externalnode.NodeClientFactory + registrationStore register.RegistrationStore + prompter terminal.Selector +} + +func defaultAllowSSHDeps() allowSSHDeps { + return allowSSHDeps{ + platform: register.LinuxPlatform{}, + nodeClients: register.DefaultNodeClientFactory{}, + registrationStore: register.NewFileRegistrationStore(), + prompter: register.TerminalPrompter{}, + } +} + +func NewCmdAllowSSH(t *terminal.Terminal, store AllowSSHStore) *cobra.Command { + cmd := &cobra.Command{ + Annotations: map[string]string{"configuration": ""}, + Use: "allow-ssh", + DisableFlagsInUseLine: true, + Short: "Trust the Brev certificate authority on this device for SSH", + Long: "Writes the Brev certificate authority to authorized_keys, allowing this device to be an SSH target for the current Linux user. Users are granted access with 'brev grant-ssh'.", + Example: " brev allow-ssh", + RunE: func(cmd *cobra.Command, args []string) error { + return runAllowSSH(cmd.Context(), t, store, defaultAllowSSHDeps()) + }, + } + + return cmd +} + +func runAllowSSH(ctx context.Context, t *terminal.Terminal, s AllowSSHStore, deps allowSSHDeps) error { + if !deps.platform.IsCompatible() { + return fmt.Errorf("brev allow-ssh is only supported on Linux") + } + + reg, err := deps.registrationStore.Load() + if err != nil { + return fmt.Errorf("failed to read registration file: %w", err) + } + + brevUser, err := s.GetCurrentUser() + if err != nil { + return fmt.Errorf("failed to get current user: %w", err) + } + + return allowSSH(ctx, t, deps, s, reg, brevUser) +} + +func allowSSH( + ctx context.Context, + t *terminal.Terminal, + deps allowSSHDeps, + tokenProvider externalnode.TokenProvider, + reg *register.DeviceRegistration, + brevUser *entity.User, +) error { + linuxUser, err := user.Current() + if err != nil { + return fmt.Errorf("failed to determine current Linux user: %w", err) + } + linuxUsername := linuxUser.Username + + checkSSHDaemon(t) + + t.Vprint("") + t.Vprint(t.Green("Allowing SSH on this device")) + t.Vprint("") + t.Vprintf(" Node: %s (%s)\n", reg.DisplayName, reg.ExternalNodeID) + t.Vprintf(" Linux user: %s\n", linuxUsername) + t.Vprint("") + + node, err := fetchRegisteredNode(ctx, deps, tokenProvider, reg) + if err != nil { + return fmt.Errorf("allow SSH failed: %w", err) + } + + if node.GetLabels()[sshcert.LabelKeySSHProvider] != sshcert.SSHProviderCertAuth { + return legacyEnableSSH(ctx, t, deps, tokenProvider, reg, brevUser, node, linuxUsername) + } + + caPublicKey := node.GetCertificateAuthority() + + if err := installCertAuthority(linuxUser, caPublicKey, reg.ExternalNodeID, linuxUsername); err != nil { + return fmt.Errorf("allow SSH failed: %w", err) + } + t.Vprint(t.Green(" Certificate authority written to authorized_keys.")) + + t.Vprint("") + t.Vprint(t.Green("SSH allowed on this device. No one has SSH access yet — grant it with: brev grant-ssh")) + return nil +} + +func legacyEnableSSH( + ctx context.Context, + t *terminal.Terminal, + deps allowSSHDeps, + tokenProvider externalnode.TokenProvider, + reg *register.DeviceRegistration, + brevUser *entity.User, + node *nodev1.ExternalNode, + linuxUsername string, +) error { + brevPortID, err := register.ResolveSSHAccessPort(ctx, t, deps.prompter, deps.nodeClients, tokenProvider, reg, node) + if err != nil { + return fmt.Errorf("allow SSH failed: %w", err) + } + + if err := register.SetupAndRegisterNodeSSHAccess(ctx, t, deps.nodeClients, tokenProvider, reg, brevUser, linuxUsername, brevPortID); err != nil { + return fmt.Errorf("allow SSH failed: %w", err) + } + + t.Vprint("") + t.Vprint(t.Green(fmt.Sprintf("SSH access enabled. You can now SSH to this device via: brev shell %s", reg.DisplayName))) + return nil +} + +func installCertAuthority(osUser *user.User, caPublicKey, nodeID, linuxUser string) error { + if caPublicKey == "" { + return fmt.Errorf("certificate authority public key is required") + } + + principal := fmt.Sprintf("brev:v1:vm:%s:login:%s", nodeID, linuxUser) + entry := fmt.Sprintf("cert-authority,principals=\"%s\" %s", principal, strings.TrimSpace(caPublicKey)) + + sshDir := filepath.Join(osUser.HomeDir, ".ssh") + if err := os.MkdirAll(sshDir, 0o700); err != nil { + return fmt.Errorf("creating .ssh directory: %w", err) + } + + authKeysPath := filepath.Join(sshDir, "authorized_keys") + + existing, err := os.ReadFile(authKeysPath) // #nosec G304 + if err != nil && !os.IsNotExist(err) { + return fmt.Errorf("reading authorized_keys: %w", err) + } + + // skip if the entry already exists. + for line := range strings.SplitSeq(string(existing), "\n") { + if strings.TrimSpace(line) == entry { + return nil + } + } + + content := string(existing) + if content != "" && !strings.HasSuffix(content, "\n") { + content += "\n" + } + content += entry + "\n" + + if err := os.WriteFile(authKeysPath, []byte(content), 0o600); err != nil { + return fmt.Errorf("writing authorized_keys: %w", err) + } + + return nil +} + +func fetchRegisteredNode( + ctx context.Context, + deps allowSSHDeps, + tokenProvider externalnode.TokenProvider, + reg *register.DeviceRegistration, +) (*nodev1.ExternalNode, error) { + client := deps.nodeClients.NewNodeClient(tokenProvider, config.GlobalConfig.GetBrevPublicAPIURL()) + resp, err := client.GetNode(ctx, connect.NewRequest(&nodev1.GetNodeRequest{ + ExternalNodeId: reg.ExternalNodeID, + })) + if err != nil { + return nil, fmt.Errorf("error retrieving node: %w", err) + } + return resp.Msg.GetExternalNode(), nil +} + +func checkSSHDaemon(t *terminal.Terminal) { + for _, svc := range []string{"ssh", "sshd"} { + out, err := exec.Command("systemctl", "is-active", svc).Output() //nolint:gosec // fixed service names + if err == nil && len(out) > 0 && string(out[:len(out)-1]) == "active" { + return + } + } + t.Vprintf(" %s\n", t.Yellow("Warning: SSH daemon does not appear to be running. SSH access may not work until sshd is started.")) +} diff --git a/pkg/cmd/enablessh/enablessh_test.go b/pkg/cmd/allowssh/allowssh_test.go similarity index 59% rename from pkg/cmd/enablessh/enablessh_test.go rename to pkg/cmd/allowssh/allowssh_test.go index 7df94144d..5fb696a48 100644 --- a/pkg/cmd/enablessh/enablessh_test.go +++ b/pkg/cmd/allowssh/allowssh_test.go @@ -1,7 +1,8 @@ -package enablessh +package allowssh import ( "context" + "fmt" "net/http/httptest" "os" "os/user" @@ -14,16 +15,16 @@ import ( "connectrpc.com/connect" "github.com/brevdev/brev-cli/pkg/cmd/register" + "github.com/brevdev/brev-cli/pkg/entity" "github.com/brevdev/brev-cli/pkg/externalnode" + "github.com/brevdev/brev-cli/pkg/terminal" ) -// tempUser returns a *user.User whose HomeDir points to a temporary directory. func tempUser(t *testing.T) *user.User { t.Helper() return &user.User{HomeDir: t.TempDir()} } -// readAuthorizedKeys is a test helper that reads ~/.ssh/authorized_keys. func readAuthorizedKeys(t *testing.T, u *user.User) string { t.Helper() data, err := os.ReadFile(filepath.Join(u.HomeDir, ".ssh", "authorized_keys")) @@ -224,17 +225,51 @@ func (m mockNodeClientFactory) NewNodeClient(provider externalnode.TokenProvider return register.NewNodeServiceClient(provider, m.serverURL) } -type mockEnableSSHStore struct { +type mockAllowSSHStore struct { token string } -func (m *mockEnableSSHStore) GetCurrentUser() (interface{}, error) { return nil, nil } -func (m *mockEnableSSHStore) GetAccessToken() (string, error) { return m.token, nil } +func (m *mockAllowSSHStore) GetCurrentUser() (*entity.User, error) { return &entity.User{}, nil } +func (m *mockAllowSSHStore) GetAccessToken() (string, error) { return m.token, nil } + +// mockSelector implements terminal.Selector, returning the first item. +type mockSelector struct{ choice string } + +func (m mockSelector) Select(_ string, items []string) string { + if m.choice != "" { + for _, s := range items { + if s == m.choice { + return s + } + } + } + if len(items) > 0 { + return items[0] + } + return "" +} -// fakeNodeService implements the server side of ExternalNodeService for testing. type fakeNodeService struct { nodev1connect.UnimplementedExternalNodeServiceHandler - getNodeFn func(*nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) + getNodeFn func(*nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) + grantCalls int + openCalls int +} + +func (f *fakeNodeService) GrantNodeSSHAccess(_ context.Context, _ *connect.Request[nodev1.GrantNodeSSHAccessRequest]) (*connect.Response[nodev1.GrantNodeSSHAccessResponse], error) { + f.grantCalls++ + return connect.NewResponse(&nodev1.GrantNodeSSHAccessResponse{}), nil +} + +func (f *fakeNodeService) OpenPort(_ context.Context, req *connect.Request[nodev1.OpenPortRequest]) (*connect.Response[nodev1.OpenPortResponse], error) { + f.openCalls++ + return connect.NewResponse(&nodev1.OpenPortResponse{ + Port: &nodev1.Port{ + PortId: fmt.Sprintf("port_%d", req.Msg.GetPortNumber()), + Protocol: req.Msg.GetProtocol(), + PortNumber: req.Msg.GetPortNumber(), + }, + }), nil } func (f *fakeNodeService) GetNode(_ context.Context, req *connect.Request[nodev1.GetNodeRequest]) (*connect.Response[nodev1.GetNodeResponse], error) { @@ -245,13 +280,14 @@ func (f *fakeNodeService) GetNode(_ context.Context, req *connect.Request[nodev1 return connect.NewResponse(resp), nil } -func startFakeServer(t *testing.T, svc *fakeNodeService) (enableSSHDeps, *httptest.Server) { +func startFakeServer(t *testing.T, svc *fakeNodeService) (allowSSHDeps, *httptest.Server) { t.Helper() _, handler := nodev1connect.NewExternalNodeServiceHandler(svc) server := httptest.NewServer(handler) t.Cleanup(server.Close) - return enableSSHDeps{ + return allowSSHDeps{ nodeClients: mockNodeClientFactory{serverURL: server.URL}, + prompter: mockSelector{}, }, server } @@ -268,7 +304,7 @@ func Test_fetchRegisteredNode(t *testing.T) { }, } deps, _ := startFakeServer(t, svc) - store := &mockEnableSSHStore{token: "tok"} + store := &mockAllowSSHStore{token: "tok"} reg := ®ister.DeviceRegistration{ExternalNodeID: "unode_abc", OrgID: "org_1"} node, err := fetchRegisteredNode(context.Background(), deps, store, reg) @@ -279,3 +315,119 @@ func Test_fetchRegisteredNode(t *testing.T) { t.Fatalf("unexpected node: %+v", node) } } + +// --- installCertAuthority --- + +func Test_installCertAuthority(t *testing.T) { + const ( + caKey = "ssh-ed25519 AAAAC3Nz dummyCA" + node = "unode_abc" + luser = "ubuntu" + ) + + t.Run("WritesLine", func(t *testing.T) { + u := tempUser(t) + if err := installCertAuthority(u, caKey, node, luser); err != nil { + t.Fatalf("installCertAuthority: %v", err) + } + want := `cert-authority,principals="brev:v1:vm:unode_abc:login:ubuntu" ssh-ed25519 AAAAC3Nz dummyCA` + if result := readAuthorizedKeys(t, u); !strings.Contains(result, want) { + t.Errorf("expected cert-authority line not found:\n%s", result) + } + }) + + t.Run("Idempotent", func(t *testing.T) { + u := tempUser(t) + for i := 0; i < 2; i++ { + if err := installCertAuthority(u, caKey, node, luser); err != nil { + t.Fatalf("installCertAuthority #%d: %v", i+1, err) + } + } + result := readAuthorizedKeys(t, u) + if n := strings.Count(result, "cert-authority"); n != 1 { + t.Errorf("expected 1 cert-authority line, got %d:\n%s", n, result) + } + }) + + t.Run("PreservesExistingKeys", func(t *testing.T) { + u := tempUser(t) + sshDir := filepath.Join(u.HomeDir, ".ssh") + if err := os.MkdirAll(sshDir, 0o700); err != nil { + t.Fatal(err) + } + original := "ssh-rsa EXISTING user@host\n" + if err := os.WriteFile(filepath.Join(sshDir, "authorized_keys"), []byte(original), 0o600); err != nil { + t.Fatal(err) + } + + if err := installCertAuthority(u, caKey, node, luser); err != nil { + t.Fatalf("installCertAuthority: %v", err) + } + + result := readAuthorizedKeys(t, u) + if !strings.Contains(result, "ssh-rsa EXISTING user@host") { + t.Errorf("existing key was removed:\n%s", result) + } + if !strings.Contains(result, "cert-authority") { + t.Errorf("cert-authority line not written:\n%s", result) + } + }) + + t.Run("EmptyKeyErrors", func(t *testing.T) { + if err := installCertAuthority(tempUser(t), "", node, luser); err == nil { + t.Error("expected error for empty CA key") + } + }) +} + +// Test_allowSSH_LegacyNodeFallsBackToKeys verifies that a node WITHOUT the +// sshprovider=certauth label skips cert-authority and uses the legacy flow: +// port resolution + key install + reflexive GrantNodeSSHAccess. +func Test_allowSSH_LegacyNodeFallsBackToKeys(t *testing.T) { + svc := &fakeNodeService{ + getNodeFn: func(_ *nodev1.GetNodeRequest) (*nodev1.GetNodeResponse, error) { + return &nodev1.GetNodeResponse{ + ExternalNode: &nodev1.ExternalNode{ + ExternalNodeId: "unode_legacy", + // No sshprovider label — legacy node. + Labels: map[string]string{}, + Ports: []*nodev1.Port{{ + PortId: "port_ssh", + Protocol: nodev1.PortProtocol_PORT_PROTOCOL_TCP, + PortNumber: 22, + }}, + }, + }, nil + }, + } + deps, _ := startFakeServer(t, svc) + + reg := ®ister.DeviceRegistration{ + DisplayName: "legacy-node", + ExternalNodeID: "unode_legacy", + OrgID: "org_1", + } + + term := terminal.New() + if err := allowSSH(context.Background(), term, deps, &mockAllowSSHStore{}, reg, &entity.User{ID: "user_1"}); err != nil { + t.Fatalf("allowSSH failed: %v", err) + } + + // Legacy flow must grant SSH access (reflexive grant). + if svc.grantCalls == 0 { + t.Error("expected GrantNodeSSHAccess to be called for legacy node") + } + + // No cert-authority line may be written for a legacy node. + realUser, err := user.Current() + if err != nil { + t.Fatalf("user.Current failed: %v", err) + } + authKeysPath := filepath.Join(realUser.HomeDir, ".ssh", "authorized_keys") + data, readErr := os.ReadFile(authKeysPath) // #nosec G304 + if readErr == nil { + if strings.Contains(string(data), "brev:v1:vm:unode_legacy") { + t.Errorf("legacy node must not write a cert-authority line:\n%s", string(data)) + } + } +} diff --git a/pkg/cmd/cmd.go b/pkg/cmd/cmd.go index 8aa4c561f..db825b7a0 100644 --- a/pkg/cmd/cmd.go +++ b/pkg/cmd/cmd.go @@ -7,6 +7,7 @@ import ( "github.com/brevdev/brev-cli/pkg/analytics" "github.com/brevdev/brev-cli/pkg/auth" "github.com/brevdev/brev-cli/pkg/cmd/agentskill" + "github.com/brevdev/brev-cli/pkg/cmd/allowssh" analyticscmd "github.com/brevdev/brev-cli/pkg/cmd/analytics" "github.com/brevdev/brev-cli/pkg/cmd/background" "github.com/brevdev/brev-cli/pkg/cmd/clipboard" @@ -15,7 +16,7 @@ import ( "github.com/brevdev/brev-cli/pkg/cmd/copy" "github.com/brevdev/brev-cli/pkg/cmd/delete" "github.com/brevdev/brev-cli/pkg/cmd/deregister" - "github.com/brevdev/brev-cli/pkg/cmd/enablessh" + "github.com/brevdev/brev-cli/pkg/cmd/disallowssh" "github.com/brevdev/brev-cli/pkg/cmd/envvars" "github.com/brevdev/brev-cli/pkg/cmd/exec" "github.com/brevdev/brev-cli/pkg/cmd/feedback" @@ -323,7 +324,8 @@ func createCmdTree(cmd *cobra.Command, t *terminal.Terminal, loginCmdStore *stor cmd.AddCommand(register.NewCmdRegister(t, externalNodeCmdStore)) cmd.AddCommand(deregister.NewCmdDeregister(t, externalNodeCmdStore)) cmd.AddCommand(upgrade.NewCmdUpgrade(t, noLoginCmdStore)) - cmd.AddCommand(enablessh.NewCmdEnableSSH(t, externalNodeCmdStore)) + cmd.AddCommand(allowssh.NewCmdAllowSSH(t, externalNodeCmdStore)) + cmd.AddCommand(disallowssh.NewCmdDisallowSSH(t, externalNodeCmdStore)) cmd.AddCommand(grantssh.NewCmdGrantSSH(t, externalNodeCmdStore)) cmd.AddCommand(revokessh.NewCmdRevokeSSH(t, externalNodeCmdStore)) cmd.AddCommand(runtasks.NewCmdRunTasks(t, noLoginCmdStore)) diff --git a/pkg/cmd/deregister/deregister.go b/pkg/cmd/deregister/deregister.go index efd9090b6..9f0526def 100644 --- a/pkg/cmd/deregister/deregister.go +++ b/pkg/cmd/deregister/deregister.go @@ -4,7 +4,10 @@ package deregister import ( "context" "fmt" + "os" "os/user" + "path/filepath" + "strings" nodev1 "buf.build/gen/go/brevdev/devplane/protocolbuffers/go/devplaneapi/v1" "connectrpc.com/connect" @@ -26,20 +29,14 @@ type DeregisterStore interface { GetAccessToken() (string, error) } -// SSHKeyRemover removes Brev-managed SSH keys and returns the lines removed. -type SSHKeyRemover interface { - RemoveBrevKeys(u *user.User) ([]string, error) +type CertAuthorityRemover interface { + RemoveCertAuthority(u *user.User, nodeID, linuxUser string) (bool, error) } -// brevSSHKeyRemover delegates to register.RemoveBrevAuthorizedKeys. -type brevSSHKeyRemover struct{} +type brevCertAuthorityRemover struct{} -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 +func (brevCertAuthorityRemover) RemoveCertAuthority(u *user.User, nodeID, linuxUser string) (bool, error) { + return removeCertAuthorityLine(u, nodeID, linuxUser) } // deregisterDeps bundles the side-effecting dependencies of runDeregister so @@ -52,7 +49,7 @@ type deregisterDeps struct { netbird register.NetBirdManager nodeClients externalnode.NodeClientFactory registrationStore register.RegistrationStore - sshKeys SSHKeyRemover + sshKeys CertAuthorityRemover } func defaultDeregisterDeps() deregisterDeps { @@ -62,9 +59,8 @@ func defaultDeregisterDeps() deregisterDeps { confirmer: register.TerminalPrompter{}, gater: sudo.Default, netbird: register.Netbird{}, - nodeClients: register.DefaultNodeClientFactory{}, + sshKeys: brevCertAuthorityRemover{}, registrationStore: register.NewFileRegistrationStore(), - sshKeys: brevSSHKeyRemover{}, } } @@ -141,7 +137,7 @@ func runDeregister(ctx context.Context, t *terminal.Terminal, s DeregisterStore, t.Vprint("") t.Vprint(t.Yellow(" This will:")) t.Vprint(" 1. Remove this node from Brev") - t.Vprint(" 2. Remove Brev SSH keys from this machine (if any)") + t.Vprint(" 2. Remove any SSH data associated with this node") t.Vprint(" 3. Uninstall the Brev tunnel") t.Vprint(" 4. Delete local registration data") t.Vprint("") @@ -168,21 +164,19 @@ func runDeregister(ctx context.Context, t *terminal.Terminal, s DeregisterStore, t.Vprintf("%s Node removed from Brev.\n", t.Green(" ✓")) t.Vprint("") - t.Vprint(t.Yellow("[Step 2/4] Removing Brev SSH keys...")) + t.Vprint(t.Yellow("[Step 2/4] Removing Brev certificate authority...")) if osUser == nil { t.Vprintf(" %s\n", t.Yellow("Skipped: could not determine current user")) } else { - removed, kerr := deps.sshKeys.RemoveBrevKeys(osUser) + linuxUsername := osUser.Username + removed, cerr := deps.sshKeys.RemoveCertAuthority(osUser, reg.ExternalNodeID, linuxUsername) switch { - case kerr != nil: - t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: failed to remove Brev SSH keys: %v", kerr))) - case len(removed) > 0: - t.Vprintf("%s Brev SSH keys removed from authorized_keys:\n", t.Green(" ✓")) - for _, key := range removed { - t.Vprintf(" - %s\n", key) - } + case cerr != nil: + t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: failed to remove cert-authority: %v", cerr))) + case removed: + t.Vprintf("%s Certificate authority removed from authorized_keys.\n", t.Green(" ✓")) default: - t.Vprint(" No Brev SSH keys found in authorized_keys.") + t.Vprint(" No certificate authority line found in authorized_keys.") } } t.Vprint("") @@ -209,3 +203,40 @@ func runDeregister(ctx context.Context, t *terminal.Terminal, s DeregisterStore, return nil } + +func removeCertAuthorityLine(osUser *user.User, nodeID, linuxUser string) (bool, error) { + principal := fmt.Sprintf("brev:v1:vm:%s:login:%s", nodeID, linuxUser) + prefix := fmt.Sprintf("cert-authority,principals=\"%s\" ", principal) + + authKeysPath := filepath.Join(osUser.HomeDir, ".ssh", "authorized_keys") + + existing, err := os.ReadFile(authKeysPath) // #nosec G304 + if err != nil { + if os.IsNotExist(err) { + return false, nil + } + return false, fmt.Errorf("reading authorized_keys: %w", err) + } + + var kept []string + var removed bool + for _, line := range strings.Split(string(existing), "\n") { + trimmed := strings.TrimSpace(line) + if strings.HasPrefix(trimmed, prefix) && strings.Contains(trimmed, "cert-authority") { + removed = true + continue + } + kept = append(kept, line) + } + + if !removed { + return false, nil + } + + result := strings.Join(kept, "\n") + if err := os.WriteFile(authKeysPath, []byte(result), 0o600); err != nil { + return false, fmt.Errorf("writing authorized_keys: %w", err) + } + + return true, nil +} diff --git a/pkg/cmd/deregister/deregister_test.go b/pkg/cmd/deregister/deregister_test.go index 95c5ac999..fe36d96e9 100644 --- a/pkg/cmd/deregister/deregister_test.go +++ b/pkg/cmd/deregister/deregister_test.go @@ -111,10 +111,10 @@ func (m mockNodeClientFactory) NewNodeClient(provider externalnode.TokenProvider type mockSSHKeyRemover struct { called bool err error - removed []string + removed bool } -func (m *mockSSHKeyRemover) RemoveBrevKeys(_ *user.User) ([]string, error) { +func (m *mockSSHKeyRemover) RemoveCertAuthority(_ *user.User, _, _ string) (bool, error) { m.called = true return m.removed, m.err } diff --git a/pkg/cmd/disallowssh/disallowssh.go b/pkg/cmd/disallowssh/disallowssh.go new file mode 100644 index 000000000..86b0b352b --- /dev/null +++ b/pkg/cmd/disallowssh/disallowssh.go @@ -0,0 +1,124 @@ +package disallowssh + +import ( + "context" + "fmt" + "os" + "os/user" + "path/filepath" + "strings" + + "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" + + "github.com/spf13/cobra" +) + +type DisallowSSHStore interface { + GetCurrentUser() (*entity.User, error) + GetAccessToken() (string, error) +} + +type disallowSSHDeps struct { + platform externalnode.PlatformChecker + registrationStore register.RegistrationStore +} + +func defaultDisallowSSHDeps() disallowSSHDeps { + return disallowSSHDeps{ + platform: register.LinuxPlatform{}, + registrationStore: register.NewFileRegistrationStore(), + } +} + +func NewCmdDisallowSSH(t *terminal.Terminal, store DisallowSSHStore) *cobra.Command { + cmd := &cobra.Command{ + Annotations: map[string]string{"configuration": ""}, + Use: "disallow-ssh", + DisableFlagsInUseLine: true, + Short: "Remove SSH certificate authority from this device", + Long: "Removes the Brev certificate authority line from authorized_keys, revoking SSH access for all users. The node remains registered.", + Example: " brev disallow-ssh", + RunE: func(cmd *cobra.Command, args []string) error { + return runDisallowSSH(cmd.Context(), t, store, defaultDisallowSSHDeps()) + }, + } + + return cmd +} + +func runDisallowSSH(_ context.Context, t *terminal.Terminal, _ DisallowSSHStore, deps disallowSSHDeps) error { + if !deps.platform.IsCompatible() { + return fmt.Errorf("brev disallow-ssh is only supported on Linux") + } + + reg, err := deps.registrationStore.Load() + if err != nil { + return fmt.Errorf("failed to read registration file: %w", err) + } + + linuxUser, err := user.Current() + if err != nil { + return fmt.Errorf("failed to determine current Linux user: %w", err) + } + + t.Vprint("") + t.Vprint(t.Green("Removing SSH certificate authority from this device")) + t.Vprint("") + t.Vprintf(" Node: %s (%s)\n", reg.DisplayName, reg.ExternalNodeID) + t.Vprintf(" Linux user: %s\n", linuxUser.Username) + t.Vprint("") + + removed, err := removeCertAuthority(linuxUser, reg.ExternalNodeID, linuxUser.Username) + if err != nil { + return fmt.Errorf("disallow SSH failed: %w", err) + } + + if removed { + t.Vprint(t.Green(" Certificate authority removed from authorized_keys.")) + } else { + t.Vprint(t.Yellow(" No certificate authority line found in authorized_keys.")) + } + + t.Vprint(t.Green("SSH disallowed. This device can no longer accept SSH certificates. Run 'brev allow-ssh' to re-enable.")) + return nil +} + +func removeCertAuthority(osUser *user.User, nodeID, linuxUser string) (bool, error) { + principal := fmt.Sprintf("brev:v1:vm:%s:login:%s", nodeID, linuxUser) + prefix := fmt.Sprintf("cert-authority,principals=\"%s\" ", principal) + + authKeysPath := filepath.Join(osUser.HomeDir, ".ssh", "authorized_keys") + + existing, err := os.ReadFile(authKeysPath) // #nosec G304 + if err != nil { + if os.IsNotExist(err) { + return false, nil + } + return false, fmt.Errorf("reading authorized_keys: %w", err) + } + + var kept []string + var removed bool + for _, line := range strings.Split(string(existing), "\n") { + trimmed := strings.TrimSpace(line) + if strings.HasPrefix(trimmed, prefix) && strings.Contains(trimmed, "cert-authority") { + removed = true + continue + } + kept = append(kept, line) + } + + if !removed { + return false, nil + } + + result := strings.Join(kept, "\n") + if err := os.WriteFile(authKeysPath, []byte(result), 0o600); err != nil { + return false, fmt.Errorf("writing authorized_keys: %w", err) + } + + return true, nil +} diff --git a/pkg/cmd/disallowssh/disallowssh_test.go b/pkg/cmd/disallowssh/disallowssh_test.go new file mode 100644 index 000000000..272a43124 --- /dev/null +++ b/pkg/cmd/disallowssh/disallowssh_test.go @@ -0,0 +1,137 @@ +package disallowssh + +import ( + "os" + "os/user" + "path/filepath" + "strings" + "testing" +) + +func tempUser(t *testing.T) *user.User { + t.Helper() + return &user.User{HomeDir: t.TempDir()} +} + +func Test_removeCertAuthority_RemovesMatchingLine(t *testing.T) { + u := tempUser(t) + sshDir := filepath.Join(u.HomeDir, ".ssh") + if err := os.MkdirAll(sshDir, 0o700); err != nil { + t.Fatal(err) + } + + caKey := "ssh-ed25519 AAAAC3Nz dummyCA" + entry := `cert-authority,principals="brev:v1:vm:unode_abc:login:ubuntu" ` + caKey + content := strings.Join([]string{ + "ssh-rsa EXISTING user@host", + entry, + "ssh-ed25519 OTHER admin@server", + "", + }, "\n") + + if err := os.WriteFile(filepath.Join(sshDir, "authorized_keys"), []byte(content), 0o600); err != nil { + t.Fatal(err) + } + + removed, err := removeCertAuthority(u, "unode_abc", "ubuntu") + if err != nil { + t.Fatalf("removeCertAuthority: %v", err) + } + if !removed { + t.Fatal("expected line to be removed") + } + + data, err := os.ReadFile(filepath.Join(sshDir, "authorized_keys")) + if err != nil { + t.Fatal(err) + } + + result := string(data) + if strings.Contains(result, caKey) { + t.Errorf("CA key still present:\n%s", result) + } + if strings.Contains(result, "cert-authority") { + t.Errorf("cert-authority line still present:\n%s", result) + } + if !strings.Contains(result, "ssh-rsa EXISTING user@host") { + t.Errorf("non-brev key was removed:\n%s", result) + } +} + +func Test_removeCertAuthority_NoopWhenFileDoesNotExist(t *testing.T) { + u := tempUser(t) + removed, err := removeCertAuthority(u, "unode_abc", "ubuntu") + if err != nil { + t.Fatalf("expected no error for missing file: %v", err) + } + if removed { + t.Error("expected removed=false for missing file") + } +} + +func Test_removeCertAuthority_NoopWhenNoMatch(t *testing.T) { + u := tempUser(t) + sshDir := filepath.Join(u.HomeDir, ".ssh") + if err := os.MkdirAll(sshDir, 0o700); err != nil { + t.Fatal(err) + } + + original := "ssh-rsa EXISTING user@host\n" + if err := os.WriteFile(filepath.Join(sshDir, "authorized_keys"), []byte(original), 0o600); err != nil { + t.Fatal(err) + } + + removed, err := removeCertAuthority(u, "unode_abc", "ubuntu") + if err != nil { + t.Fatalf("removeCertAuthority: %v", err) + } + if removed { + t.Error("expected removed=false when no match") + } + + data, err := os.ReadFile(filepath.Join(sshDir, "authorized_keys")) + if err != nil { + t.Fatal(err) + } + if string(data) != original { + t.Errorf("file was modified when it shouldn't have been") + } +} + +func Test_removeCertAuthority_OnlyRemovesMatchingPrincipal(t *testing.T) { + u := tempUser(t) + sshDir := filepath.Join(u.HomeDir, ".ssh") + if err := os.MkdirAll(sshDir, 0o700); err != nil { + t.Fatal(err) + } + + otherEntry := `cert-authority,principals="brev:v1:vm:other_node:login:ubuntu" ssh-ed25519 OTHER_CA` + targetEntry := `cert-authority,principals="brev:v1:vm:unode_abc:login:ubuntu" ssh-ed25519 TARGET_CA` + content := strings.Join([]string{ + otherEntry, + targetEntry, + "", + }, "\n") + + if err := os.WriteFile(filepath.Join(sshDir, "authorized_keys"), []byte(content), 0o600); err != nil { + t.Fatal(err) + } + + removed, err := removeCertAuthority(u, "unode_abc", "ubuntu") + if err != nil { + t.Fatalf("removeCertAuthority: %v", err) + } + if !removed { + t.Fatal("expected line to be removed") + } + + data, _ := os.ReadFile(filepath.Join(sshDir, "authorized_keys")) + result := string(data) + + if strings.Contains(result, "TARGET_CA") { + t.Errorf("target CA still present:\n%s", result) + } + if !strings.Contains(result, "OTHER_CA") { + t.Errorf("other node's CA was removed:\n%s", result) + } +} diff --git a/pkg/cmd/enablessh/enablessh.go b/pkg/cmd/enablessh/enablessh.go deleted file mode 100644 index 9788b0e6f..000000000 --- a/pkg/cmd/enablessh/enablessh.go +++ /dev/null @@ -1,153 +0,0 @@ -// Package enablessh provides the brev enableSSH command for enabling SSH access -// to a registered external node. -package enablessh - -import ( - "context" - "fmt" - "os/exec" - "os/user" - - nodev1 "buf.build/gen/go/brevdev/devplane/protocolbuffers/go/devplaneapi/v1" - "connectrpc.com/connect" - - "github.com/brevdev/brev-cli/pkg/cmd/register" - "github.com/brevdev/brev-cli/pkg/config" - "github.com/brevdev/brev-cli/pkg/entity" - breverrors "github.com/brevdev/brev-cli/pkg/errors" - "github.com/brevdev/brev-cli/pkg/externalnode" - "github.com/brevdev/brev-cli/pkg/terminal" - - "github.com/spf13/cobra" -) - -// EnableSSHStore defines the store methods needed by the enableSSH command. -type EnableSSHStore interface { - GetCurrentUser() (*entity.User, error) - GetAccessToken() (string, error) -} - -// enableSSHDeps bundles the side-effecting dependencies of runEnableSSH so they -// can be replaced in tests. -type enableSSHDeps struct { - platform externalnode.PlatformChecker - nodeClients externalnode.NodeClientFactory - registrationStore register.RegistrationStore - prompter terminal.Selector -} - -func defaultEnableSSHDeps() enableSSHDeps { - return enableSSHDeps{ - platform: register.LinuxPlatform{}, - nodeClients: register.DefaultNodeClientFactory{}, - registrationStore: register.NewFileRegistrationStore(), - prompter: register.TerminalPrompter{}, - } -} - -func NewCmdEnableSSH(t *terminal.Terminal, store EnableSSHStore) *cobra.Command { - cmd := &cobra.Command{ - Annotations: map[string]string{"configuration": ""}, - Use: "enable-ssh", - DisableFlagsInUseLine: true, - Short: "Enable SSH access to this registered device", - Long: "Enable SSH access to this registered device for the current Brev user.", - Example: " brev enable-ssh", - RunE: func(cmd *cobra.Command, args []string) error { - return runEnableSSH(cmd.Context(), t, store, defaultEnableSSHDeps()) - }, - } - - return cmd -} - -func runEnableSSH(ctx context.Context, t *terminal.Terminal, s EnableSSHStore, deps enableSSHDeps) error { - if !deps.platform.IsCompatible() { - return fmt.Errorf("brev enable-ssh is only supported on Linux") - } - - reg, err := deps.registrationStore.Load() - if err != nil { - return fmt.Errorf("failed to read registration file: %w", err) - } - - brevUser, err := s.GetCurrentUser() - if err != nil { - return breverrors.WrapAndTrace(err) - } - - return enableSSH(ctx, t, deps, s, reg, brevUser) -} - -// enableSSH grants SSH access to the given node for the current Brev user. -// This is the "reflexive grant" — granting yourself SSH access to the device. -func enableSSH( - ctx context.Context, - t *terminal.Terminal, - deps enableSSHDeps, - tokenProvider externalnode.TokenProvider, - reg *register.DeviceRegistration, - brevUser *entity.User, -) error { - linuxUser, err := user.Current() - if err != nil { - return fmt.Errorf("failed to determine current Linux user: %w", err) - } - linuxUsername := linuxUser.Username - - checkSSHDaemon(t) - - t.Vprint("") - t.Vprint(t.Green("Enabling SSH access on this device")) - t.Vprint("") - t.Vprintf(" Node: %s (%s)\n", reg.DisplayName, reg.ExternalNodeID) - t.Vprintf(" Brev user: %s\n", brevUser.ID) - t.Vprintf(" Linux user: %s\n", linuxUsername) - t.Vprint("") - - node, err := fetchRegisteredNode(ctx, deps, tokenProvider, reg) - if err != nil { - return fmt.Errorf("enable SSH failed: %w", err) - } - - brevPortID, err := register.ResolveSSHAccessPort(ctx, t, deps.prompter, deps.nodeClients, tokenProvider, reg, node) - if err != nil { - return fmt.Errorf("enable SSH failed: %w", err) - } - - if err := register.SetupAndRegisterNodeSSHAccess(ctx, t, deps.nodeClients, tokenProvider, reg, brevUser, linuxUsername, brevPortID); err != nil { - return fmt.Errorf("enable SSH failed: %w", err) - } - - t.Vprint(t.Green(fmt.Sprintf("SSH access enabled. You can now SSH to this device via: brev shell %s", reg.DisplayName))) - return nil -} - -func fetchRegisteredNode( - ctx context.Context, - deps enableSSHDeps, - tokenProvider externalnode.TokenProvider, - reg *register.DeviceRegistration, -) (*nodev1.ExternalNode, error) { - client := deps.nodeClients.NewNodeClient(tokenProvider, config.GlobalConfig.GetBrevPublicAPIURL()) - resp, err := client.GetNode(ctx, connect.NewRequest(&nodev1.GetNodeRequest{ - ExternalNodeId: reg.ExternalNodeID, - OrganizationId: reg.OrgID, - })) - if err != nil { - return nil, fmt.Errorf("error retrieving node: %w", err) - } - return resp.Msg.GetExternalNode(), nil -} - -// checkSSHDaemon prints a warning if neither "ssh" nor "sshd" systemd services -// appear to be active. It never returns an error — it is best-effort. -func checkSSHDaemon(t *terminal.Terminal) { - for _, svc := range []string{"ssh", "sshd"} { - out, err := exec.Command("systemctl", "is-active", svc).Output() //nolint:gosec // fixed service names - if err == nil && len(out) > 0 && string(out[:len(out)-1]) == "active" { - return - } - } - t.Vprintf(" %s\n", t.Yellow("Warning: SSH daemon does not appear to be running. SSH access may not work until sshd is started.")) -} diff --git a/pkg/cmd/grantssh/grantssh.go b/pkg/cmd/grantssh/grantssh.go index 35985834e..05ab0b01a 100644 --- a/pkg/cmd/grantssh/grantssh.go +++ b/pkg/cmd/grantssh/grantssh.go @@ -187,7 +187,7 @@ func runGrantSSH(ctx context.Context, t *terminal.Terminal, s GrantSSHStore, opt t.Vprint("") linuxUser = deps.prompter.Select("Select Linux user on the node", linuxUserOptions) } else { - return fmt.Errorf("no Linux users on this node yet; run with --linux-user to specify one (e.g. after enable-ssh on the node)") + return fmt.Errorf("no Linux users on this node yet; run with --linux-user to specify one (e.g. after allow-ssh on the node)") } } else { selectedUser, err = findUserByIDOrEmail(orgMembers, opts.userIDOrEmail) diff --git a/pkg/cmd/register/register.go b/pkg/cmd/register/register.go index 2ad88b435..42fa62e96 100644 --- a/pkg/cmd/register/register.go +++ b/pkg/cmd/register/register.go @@ -5,7 +5,6 @@ import ( "context" "errors" "fmt" - "os/user" "strings" "time" @@ -81,7 +80,9 @@ func defaultRegisterDeps() registerDeps { var ( registerLong = `Register your device with NVIDIA Brev -This command sets up network connectivity and registers this machine with Brev. +This command registers this machine with Brev and brings up the Brev tunnel. +Registration does not enable SSH; run 'brev allow-ssh' afterwards to allow SSH +on this device, then 'brev grant-ssh' to grant users SSH access. Two modes are supported: • Interactive (default): run 'brev register' with no flags and follow prompts for device name, org, and options. @@ -92,7 +93,10 @@ Two modes are supported: # 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` + + + # Enable SSH access to this device after registering + brev allow-ssh` ) func NewCmdRegister(t *terminal.Terminal, store RegisterStore) *cobra.Command { @@ -155,8 +159,8 @@ func runRegister(ctx context.Context, t *terminal.Terminal, s RegisterStore, opt } } - // Run through the login flow - brevUser, err := s.GetCurrentUser() + // Verify the user is authenticated before any local side effects. + _, err := s.GetCurrentUser() if err != nil { return breverrors.WrapAndTrace(err) } @@ -227,35 +231,12 @@ func runRegister(ctx context.Context, t *terminal.Terminal, s RegisterStore, opt } // Perform the registration steps - reg, err := runRegisterSteps(ctx, t, s, name, org, deps) + _, err = runRegisterSteps(ctx, t, s, name, org, deps) if err != nil { return err } - // Determine if SSH access should be enabled - enableSSH := false - sshPortForGrant := int32(0) - if opts.interactive { - enableSSH = deps.prompter.ConfirmYesNo("Would you like to enable SSH access to this device?") - if enableSSH { - sshPortForGrant = 0 // prompt for port - } - } else if opts.sshPort != 0 { - enableSSH = true - sshPortForGrant = opts.sshPort - } - - // Grant SSH access if requested - if enableSSH { - osUser, err := user.Current() - if err != nil { - return fmt.Errorf("failed to determine current Linux user: %w", err) - } - if err := grantSSHAccessWithPort(ctx, t, deps, s, reg, brevUser, osUser, sshPortForGrant, opts.interactive, opts.skipConfirm); err != nil { - t.Vprintf(" %s\n", t.Yellow(fmt.Sprintf("Warning: %v", err))) - } - } - + suggestAllowSSH(t) return nil } @@ -291,6 +272,7 @@ func runRegisterSteps(ctx context.Context, t *terminal.Terminal, s RegisterStore Name: name, DeviceId: deviceID, NodeSpec: toProtoNodeSpec(hwProfile), + Labels: map[string]string{"sshprovider": "certauth"}, })) if err != nil { // dev-plane returns CodeAlreadyExists for a duplicate node name; surface @@ -424,58 +406,7 @@ 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 - } - +func suggestAllowSSH(t *terminal.Terminal) { 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 + t.Vprintf(" %s\n", t.Green("To allow SSH on this device, run: brev allow-ssh")) } diff --git a/pkg/cmd/register/register_test.go b/pkg/cmd/register/register_test.go index d98b1a92a..4ebb394bf 100644 --- a/pkg/cmd/register/register_test.go +++ b/pkg/cmd/register/register_test.go @@ -662,109 +662,6 @@ 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) { tests := []struct { name string @@ -979,153 +876,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) - }) - } -}