Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions pkg/cmd/cmd.go
Original file line number Diff line number Diff line change
Expand Up @@ -265,12 +265,14 @@ func NewBrevCommand() *cobra.Command { //nolint:funlen,gocognit,gocyclo // defin
fmt.Printf("%v\n", err)
}

createCmdTree(cmds, t, loginCmdStore, noLoginCmdStore, loginAuth, externalNodeCmdStore)
accessKeyAuth := register.NewAccessKeyAuthenticator(memAuthStore.MemoryAuthStore, loginAuth)

createCmdTree(cmds, t, loginCmdStore, noLoginCmdStore, loginAuth, externalNodeCmdStore, accessKeyAuth)

return cmds
}

func createCmdTree(cmd *cobra.Command, t *terminal.Terminal, loginCmdStore *store.AuthHTTPStore, noLoginCmdStore *store.AuthHTTPStore, loginAuth *auth.LoginAuth, externalNodeCmdStore *store.AuthHTTPStore) { //nolint:funlen // define brev command
func createCmdTree(cmd *cobra.Command, t *terminal.Terminal, loginCmdStore *store.AuthHTTPStore, noLoginCmdStore *store.AuthHTTPStore, loginAuth *auth.LoginAuth, externalNodeCmdStore *store.AuthHTTPStore, accessKeyAuth register.AccessKeyAuthenticator) { //nolint:funlen // define brev command
cmd.AddCommand(set.NewCmdSet(t, loginCmdStore, noLoginCmdStore))
cmd.AddCommand(ls.NewCmdLs(t, loginCmdStore, noLoginCmdStore))
cmd.AddCommand(org.NewCmdOrg(t, loginCmdStore, noLoginCmdStore))
Expand Down Expand Up @@ -318,7 +320,7 @@ 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(register.NewCmdRegister(t, externalNodeCmdStore, accessKeyAuth))
cmd.AddCommand(deregister.NewCmdDeregister(t, externalNodeCmdStore))
cmd.AddCommand(upgrade.NewCmdUpgrade(t, noLoginCmdStore))
cmd.AddCommand(enablessh.NewCmdEnableSSH(t, externalNodeCmdStore))
Expand Down
33 changes: 25 additions & 8 deletions pkg/cmd/deregister/deregister.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package deregister

import (
"context"
"errors"
"fmt"
"os/user"

Expand Down Expand Up @@ -106,7 +107,7 @@ func runDeregister(ctx context.Context, t *terminal.Terminal, s DeregisterStore,
return fmt.Errorf("sudo issue: %w", err)
}

reg, err := deps.registrationStore.Load()
reg, err := deps.registrationStore.Load(true) // deregister should still work for pending registrations
if err != nil {
return err //nolint:wrapcheck // do not present stack trace for this error
}
Expand Down Expand Up @@ -158,14 +159,30 @@ func runDeregister(ctx context.Context, t *terminal.Terminal, s DeregisterStore,
}

t.Vprint(t.Yellow("[Step 1/4] Removing node from Brev..."))
client := deps.nodeClients.NewNodeClient(s, config.GlobalConfig.GetBrevPublicAPIURL())
_, err = client.RemoveNode(ctx, connect.NewRequest(&nodev1.RemoveNodeRequest{
ExternalNodeId: reg.ExternalNodeID,
}))
if err != nil {
return fmt.Errorf("failed to deregister node: %w", err)
if reg.ExternalNodeID == "" {
// Pending registration: AddNode was never confirmed, so the CLI has no
// external node ID to remove. Skip RemoveNode and clean up local state
// so the user isn't stuck with an undeletable pending record.
t.Vprintf(" %s\n", t.Yellow("No registered node to remove (pending registration); cleaning up local state."))
} else {
client := deps.nodeClients.NewNodeClient(s, config.GlobalConfig.GetBrevPublicAPIURL())
_, err = client.RemoveNode(ctx, connect.NewRequest(&nodev1.RemoveNodeRequest{
ExternalNodeId: reg.ExternalNodeID,
}))
if err != nil {
// NotFound means the node is already gone (e.g. a previous attempt
// succeeded before a transient error). Treat as success so local
// cleanup proceeds and a retry isn't blocked.
var connectErr *connect.Error
if errors.As(err, &connectErr) && connectErr.Code() == connect.CodeNotFound {
t.Vprintf(" %s\n", t.Yellow("Node not found on Brev (already removed); continuing."))
} else {
return fmt.Errorf("failed to deregister node: %w", err)
}
} else {
t.Vprintf("%s Node removed from Brev.\n", t.Green(" ✓"))
}
}
t.Vprintf("%s Node removed from Brev.\n", t.Green(" ✓"))
t.Vprint("")

t.Vprint(t.Yellow("[Step 2/4] Removing Brev SSH keys..."))
Expand Down
86 changes: 84 additions & 2 deletions pkg/cmd/deregister/deregister_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -57,7 +57,7 @@ func (m *mockRegistrationStore) Save(reg *register.DeviceRegistration) error {
return nil
}

func (m *mockRegistrationStore) Load() (*register.DeviceRegistration, error) {
func (m *mockRegistrationStore) Load(bool) (*register.DeviceRegistration, error) {
if m.reg == nil {
return nil, fmt.Errorf("no registration")
}
Expand Down Expand Up @@ -281,7 +281,6 @@ func Test_runDeregister_RemoveNodeFails(t *testing.T) {
t.Fatal("expected error when RemoveNode fails")
}

// Registration should still exist (server-side removal failed)
exists, err := regStore.Exists()
if err != nil {
t.Fatalf("Exists error: %v", err)
Expand All @@ -291,6 +290,89 @@ func Test_runDeregister_RemoveNodeFails(t *testing.T) {
}
}

func Test_runDeregister_RemoveNodeNotFound_ProceedsCleanup(t *testing.T) {
regStore := &mockRegistrationStore{
reg: &register.DeviceRegistration{
ExternalNodeID: "unode_abc",
DisplayName: "My Spark",
OrgID: "org_123",
},
}

store := &mockDeregisterStore{
user: &entity.User{ID: "user_1"},
token: "tok",
}

svc := &fakeNodeService{
removeNodeFn: func(_ *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) {
return nil, connect.NewError(connect.CodeNotFound, nil)
},
}

deps, server := testDeregisterDeps(t, svc, regStore)
defer server.Close()

term := terminal.New()
err := runDeregister(context.Background(), term, store, deps, false)
if err != nil {
t.Fatalf("NotFound should be treated as success (node already gone), got: %v", err)
}

exists, err := regStore.Exists()
if err != nil {
t.Fatalf("Exists error: %v", err)
}
if exists {
t.Error("expected local registration to be deleted even when RemoveNode returns NotFound")
}
}

func Test_runDeregister_PendingRegistration_SkipsRemoveNodeAndCleansUp(t *testing.T) {
regStore := &mockRegistrationStore{
reg: &register.DeviceRegistration{
DisplayName: "My Spark",
OrgID: "org_123",
DeviceID: "dev-uuid-pending",
Status: register.RegistrationStatusPending,
},
}

store := &mockDeregisterStore{
user: &entity.User{ID: "user_1"},
token: "tok",
}

var removeCalled bool
svc := &fakeNodeService{
removeNodeFn: func(req *nodev1.RemoveNodeRequest) (*nodev1.RemoveNodeResponse, error) {
removeCalled = true
return nil, fmt.Errorf("RemoveNode should not be called with empty ID %q", req.GetExternalNodeId())
},
}

deps, server := testDeregisterDeps(t, svc, regStore)
defer server.Close()

term := terminal.New()
err := runDeregister(context.Background(), term, store, deps, false)
if err != nil {
t.Fatalf("deregister of a pending registration should succeed, got: %v", err)
}

if removeCalled {
t.Error("RemoveNode should not be called for a pending registration (no ExternalNodeID)")
}

exists, err := regStore.Exists()
if err != nil {
t.Fatalf("Exists error: %v", err)
}
if exists {
t.Error("expected local pending registration to be deleted")
}
}

func Test_runDeregister_AlwaysUninstallsNetbird(t *testing.T) {
regStore := &mockRegistrationStore{
reg: &register.DeviceRegistration{
Expand Down
2 changes: 1 addition & 1 deletion pkg/cmd/enablessh/enablessh.go
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,7 @@ func runEnableSSH(ctx context.Context, t *terminal.Terminal, s EnableSSHStore, d
return fmt.Errorf("brev enable-ssh is only supported on Linux")
}

reg, err := deps.registrationStore.Load()
reg, err := deps.registrationStore.Load(false)
if err != nil {
return fmt.Errorf("failed to read registration file: %w", err)
}
Expand Down
2 changes: 1 addition & 1 deletion pkg/cmd/grantssh/grantssh_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,7 @@ func (m *mockRegistrationStore) Save(reg *register.DeviceRegistration) error {
return nil
}

func (m *mockRegistrationStore) Load() (*register.DeviceRegistration, error) {
func (m *mockRegistrationStore) Load(bool) (*register.DeviceRegistration, error) {
if m.reg == nil {
return nil, fmt.Errorf("no registration")
}
Expand Down
38 changes: 25 additions & 13 deletions pkg/cmd/register/device_registration_store.go
Original file line number Diff line number Diff line change
Expand Up @@ -20,31 +20,35 @@ const (
globalRegistrationDir = "/etc/brev"
)

const (
RegistrationStatusPending = "pending" // used for retries
RegistrationStatusRegistered = "registered"
)

// DeviceRegistration is the persistent identity file for a registered device.
// Fields align with the AddNodeResponse from dev-plane.
type DeviceRegistration struct {
ExternalNodeID string `json:"external_node_id"`
DisplayName string `json:"display_name"`
OrgID string `json:"org_id"`
OrgName string `json:"org_name"`
DeviceID string `json:"device_id"`
RegisteredAt string `json:"registered_at"`
HardwareProfile HardwareProfile `json:"hardware_profile"`
ExternalNodeID string `json:"external_node_id"`
DisplayName string `json:"display_name"`
OrgID string `json:"org_id"`
OrgName string `json:"org_name"`
DeviceID string `json:"device_id"`
RegisteredAt string `json:"registered_at"`
HardwareProfile HardwareProfile `json:"hardware_profile"`
RegistrationToken string `json:"registration_token,omitempty"`
Status string `json:"status,omitempty"`
}

// RegistrationStore defines the contract for persisting device registration data.
type RegistrationStore interface {
Save(reg *DeviceRegistration) error
Load() (*DeviceRegistration, error)
Load(includeAll bool) (*DeviceRegistration, error)
Delete() error
Exists() (bool, error)
}

// FileRegistrationStore implements RegistrationStore using the global /etc/brev/ path.
type FileRegistrationStore struct{}

// NewFileRegistrationStore returns a FileRegistrationStore that reads/writes
// from /etc/brev/device_registration.json.
func NewFileRegistrationStore() *FileRegistrationStore {
return &FileRegistrationStore{}
}
Expand Down Expand Up @@ -72,8 +76,7 @@ func (s *FileRegistrationStore) Save(reg *DeviceRegistration) error {
return sudoWriteFile(path, data)
}

// Load reads the registration file and returns the parsed DeviceRegistration
func (s *FileRegistrationStore) Load() (*DeviceRegistration, error) {
func (s *FileRegistrationStore) Load(includeAll bool) (*DeviceRegistration, error) {
path := s.path()
exists, err := s.Exists()
if !exists {
Expand All @@ -86,7 +89,16 @@ func (s *FileRegistrationStore) Load() (*DeviceRegistration, error) {
if err := files.ReadJSON(files.AppFs, path, &reg); err != nil {
return nil, breverrors.WrapAndTrace(err)
}
if includeAll {
if reg.OrgID == "" && reg.DeviceID == "" {
return nil, breverrors.New("malformed registration")
}
return &reg, nil
}
if reg.ExternalNodeID == "" || reg.OrgID == "" {
if reg.Status == RegistrationStatusPending {
return nil, breverrors.New("device registration is incomplete; re-run 'brev register' to finish")
}
return nil, breverrors.New("malformed registration")
}
return &reg, nil
Expand Down
69 changes: 65 additions & 4 deletions pkg/cmd/register/device_registration_store_test.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package register

import (
"strings"
"testing"

"github.com/brevdev/brev-cli/pkg/files"
Expand Down Expand Up @@ -42,7 +43,7 @@ func Test_SaveAndLoadRegistration_RoundTrip(t *testing.T) {
t.Fatalf("Save failed: %v", err)
}

loaded, err := store.Load()
loaded, err := store.Load(false)
if err != nil {
t.Fatalf("Load failed: %v", err)
}
Expand Down Expand Up @@ -138,7 +139,7 @@ func Test_LoadRegistration_FailsWhenMissing(t *testing.T) {

store := NewFileRegistrationStore()

_, err := store.Load()
_, err := store.Load(false)
if err == nil {
t.Error("expected error loading missing registration")
}
Expand All @@ -159,7 +160,7 @@ func Test_LoadRegistration_RejectsMissingExternalNodeID(t *testing.T) {
t.Fatalf("Save failed: %v", err)
}

_, err := store.Load()
_, err := store.Load(false)
if err == nil {
t.Fatal("expected error loading registration with empty ExternalNodeID")
}
Expand All @@ -180,7 +181,7 @@ func Test_LoadRegistration_RejectsMissingOrgID(t *testing.T) {
t.Fatalf("Save failed: %v", err)
}

_, err := store.Load()
_, err := store.Load(false)
if err == nil {
t.Fatal("expected error loading registration with empty OrgID")
}
Expand All @@ -197,3 +198,63 @@ func Test_DeleteRegistration_FailsWhenMissing(t *testing.T) {
t.Error("expected error deleting missing registration")
}
}

func Test_Load_IncludeAllReturnsPendingRecord(t *testing.T) {
cleanup := setupTestFs(t)
defer cleanup()

store := NewFileRegistrationStore()

pending := &DeviceRegistration{
DisplayName: "My Spark",
OrgID: "org_xyz",
DeviceID: "device-uuid-123",
Status: RegistrationStatusPending,
}
if err := store.Save(pending); err != nil {
t.Fatalf("Save failed: %v", err)
}

loaded, err := store.Load(true)
if err != nil {
t.Fatalf("Load failed: %v", err)
}
if loaded.DeviceID != "device-uuid-123" {
t.Errorf("DeviceID mismatch: got %s, want device-uuid-123", loaded.DeviceID)
}
if loaded.Status != RegistrationStatusPending {
t.Errorf("Status mismatch: got %q, want %q", loaded.Status, RegistrationStatusPending)
}
if loaded.ExternalNodeID != "" {
t.Errorf("pending record should have no ExternalNodeID, got %q", loaded.ExternalNodeID)
}

if _, err := store.Load(false); err == nil {
t.Error("expected Load(false) to error on a pending record")
}
}

func Test_Load_PendingRecordErrorMessage(t *testing.T) {
cleanup := setupTestFs(t)
defer cleanup()

store := NewFileRegistrationStore()

pending := &DeviceRegistration{
DisplayName: "My Spark",
OrgID: "org_xyz",
DeviceID: "device-uuid-123",
Status: RegistrationStatusPending,
}
if err := store.Save(pending); err != nil {
t.Fatalf("Save failed: %v", err)
}

_, err := store.Load(false)
if err == nil {
t.Fatal("expected Load() to error on a pending record")
}
if !strings.Contains(err.Error(), "incomplete") {
t.Errorf("expected 'incomplete' in error, got: %v", err)
}
}
Loading
Loading