diff --git a/.goreleaser_ui.yaml b/.goreleaser_ui.yaml index 1157e6379..ca5148823 100644 --- a/.goreleaser_ui.yaml +++ b/.goreleaser_ui.yaml @@ -93,7 +93,9 @@ nfpms: - src: client/ui/build/appicon.png dst: /usr/share/pixmaps/netbird.png dependencies: - - netbird + - netbird (>= 0.75.0) + - libgtk-4-1 (>= 4.14) + - libwebkitgtk-6.0-4 - maintainer: Netbird description: Netbird client UI. @@ -114,7 +116,9 @@ nfpms: - src: client/ui/build/appicon.png dst: /usr/share/pixmaps/netbird.png dependencies: - - netbird + - netbird >= 0.75.0 + - (gtk4 >= 4.14 or libgtk-4-1 >= 4.14) + - (webkitgtk6.0 or libwebkitgtk-6_0-4) rpm: signature: diff --git a/client/internal/connect.go b/client/internal/connect.go index 87126b222..ceb39419e 100644 --- a/client/internal/connect.go +++ b/client/internal/connect.go @@ -113,11 +113,14 @@ func (c *ConnectClient) RunOnAndroid( stateFilePath string, cacheDir string, ) error { + notifier := tunnelnotifier.New(networkChangeListener, nil) + defer notifier.Close() + // in case of non Android os these variables will be nil mobileDependency := MobileDependency{ TunAdapter: tunAdapter, IFaceDiscover: iFaceDiscover, - NetworkChangeListener: networkChangeListener, + NetworkChangeListener: notifier, HostDNSAddresses: dnsAddresses, DnsReadyListener: dnsReadyListener, StateFilePath: stateFilePath, diff --git a/client/internal/dns/interface_index.go b/client/internal/dns/interface_index.go new file mode 100644 index 000000000..9e7dca080 --- /dev/null +++ b/client/internal/dns/interface_index.go @@ -0,0 +1,15 @@ +package dns + +import ( + "fmt" + "net" +) + +func getInterfaceIndex(interfaceName string) (int, error) { + iface, err := net.InterfaceByName(interfaceName) + if err != nil { + return 0, fmt.Errorf("lookup interface %q: %w", interfaceName, err) + } + + return iface.Index, nil +} diff --git a/client/internal/dns/interface_index_test.go b/client/internal/dns/interface_index_test.go new file mode 100644 index 000000000..9b146398a --- /dev/null +++ b/client/internal/dns/interface_index_test.go @@ -0,0 +1,35 @@ +package dns + +import ( + "net" + "testing" +) + +func TestGetInterfaceIndexExisting(t *testing.T) { + interfaces, err := net.Interfaces() + if err != nil { + t.Fatalf("list network interfaces: %v", err) + } + if len(interfaces) == 0 { + t.Fatal("expected at least one network interface") + } + + iface := interfaces[0] + index, err := getInterfaceIndex(iface.Name) + if err != nil { + t.Fatalf("look up existing interface %q: %v", iface.Name, err) + } + if index != iface.Index { + t.Fatalf("expected interface index %d, got %d", iface.Index, index) + } +} + +func TestGetInterfaceIndexMissing(t *testing.T) { + index, err := getInterfaceIndex("netbird-interface-that-does-not-exist") + if index != 0 { + t.Fatalf("expected missing interface index to be 0, got %d", index) + } + if err == nil { + t.Fatal("expected missing interface lookup to return an error") + } +} diff --git a/client/internal/dns/notifier.go b/client/internal/dns/notifier.go index 35cb6ff82..79d924a78 100644 --- a/client/internal/dns/notifier.go +++ b/client/internal/dns/notifier.go @@ -51,7 +51,5 @@ func (n *notifier) notify() { return } - go func(l listener.NetworkChangeListener) { - l.OnNetworkChanged("") - }(n.listener) + n.listener.OnNetworkChanged("") } diff --git a/client/internal/dns/upstream_ios.go b/client/internal/dns/upstream_ios.go index b989bf0f9..793d87fca 100644 --- a/client/internal/dns/upstream_ios.go +++ b/client/internal/dns/upstream_ios.go @@ -130,8 +130,3 @@ func GetClientPrivate(iface privateClientIface, upstreamIP netip.Addr, dialTimeo } return client, nil } - -func getInterfaceIndex(interfaceName string) (int, error) { - iface, err := net.InterfaceByName(interfaceName) - return iface.Index, err -} diff --git a/client/internal/profilemanager/state.go b/client/internal/profilemanager/state.go index 9e9577796..fcd1c384c 100644 --- a/client/internal/profilemanager/state.go +++ b/client/internal/profilemanager/state.go @@ -45,12 +45,35 @@ func (pm *ProfileManager) GetProfileState(id ID) (*ProfileState, error) { return &state, nil } -func (pm *ProfileManager) SetActiveProfileState(state *ProfileState) error { +// SetProfileState writes the state file of the profile identified by id. Prefer +// it over SetActiveProfileState whenever the caller knows which profile the data +// belongs to: an SSO login spans seconds of user interaction, and the active +// profile can change during it, which would file the account email under +// whichever profile happened to be active when the flow returned. +func (pm *ProfileManager) SetProfileState(id ID, state *ProfileState) error { configDir, err := getConfigDir() if err != nil { return fmt.Errorf("get config directory: %w", err) } + if id == "" { + return fmt.Errorf("empty profile ID") + } + if id != defaultProfileName && !IsValidProfileFilenameStem(id) { + return fmt.Errorf("invalid profile ID: %q", id) + } + + stateFile := filepath.Join(configDir, id.String()+".state.json") + if err := util.WriteJsonWithRestrictedPermission(context.Background(), stateFile, state); err != nil { + return fmt.Errorf("write profile state: %w", err) + } + + return nil +} + +// SetActiveProfileState writes the state file of whichever profile is active at +// call time. Use SetProfileState when the target profile is known. +func (pm *ProfileManager) SetActiveProfileState(state *ProfileState) error { activeProf, err := pm.GetActiveProfile() if err != nil { if errors.Is(err, ErrNoActiveProfile) { @@ -59,18 +82,7 @@ func (pm *ProfileManager) SetActiveProfileState(state *ProfileState) error { return fmt.Errorf("get active profile: %w", err) } - id := activeProf.ID - if id != defaultProfileName && !IsValidProfileFilenameStem(id) { - return fmt.Errorf("invalid active profile ID: %q", id) - } - - stateFile := filepath.Join(configDir, id.String()+".state.json") - err = util.WriteJsonWithRestrictedPermission(context.Background(), stateFile, state) - if err != nil { - return fmt.Errorf("write profile state: %w", err) - } - - return nil + return pm.SetProfileState(activeProf.ID, state) } // RemoveProfileState deletes the per-profile state file (which holds the diff --git a/client/internal/routemanager/dnsinterceptor/handler.go b/client/internal/routemanager/dnsinterceptor/handler.go index f92300bfd..d20f4b944 100644 --- a/client/internal/routemanager/dnsinterceptor/handler.go +++ b/client/internal/routemanager/dnsinterceptor/handler.go @@ -479,7 +479,7 @@ func (d *DnsInterceptor) removeDNATMappings(realPrefixes []netip.Prefix, logger // internalDnatFw checks if the firewall supports internal DNAT func (d *DnsInterceptor) internalDnatFw() (internalDNATer, bool) { - if d.firewall == nil || runtime.GOOS != "android" { + if d.firewall == nil || d.fakeIPManager == nil || runtime.GOOS != "android" { return nil, false } fw, ok := d.firewall.(internalDNATer) diff --git a/client/internal/routemanager/manager.go b/client/internal/routemanager/manager.go index 2ab7e2a85..0cb74fd45 100644 --- a/client/internal/routemanager/manager.go +++ b/client/internal/routemanager/manager.go @@ -165,31 +165,36 @@ func (m *DefaultManager) setupAndroidRoutes(config ManagerConfig) { routesForComparison := slices.Clone(cr) if config.DNSFeatureFlag { - m.fakeIPManager = fakeip.NewManager() - - v4ID := uuid.NewString() - fakeIPRoute := &route.Route{ - ID: route.ID(v4ID), - Network: m.fakeIPManager.GetFakeIPBlock(), - NetID: route.NetID(v4ID), - Peer: m.pubKey, - NetworkType: route.IPv4Network, - } - v6ID := uuid.NewString() - fakeIPv6Route := &route.Route{ - ID: route.ID(v6ID), - Network: m.fakeIPManager.GetFakeIPv6Block(), - NetID: route.NetID(v6ID), - Peer: m.pubKey, - NetworkType: route.IPv6Network, - } - cr = append(cr, fakeIPRoute, fakeIPv6Route) - m.notifier.SetFakeIPRoutes([]*route.Route{fakeIPRoute, fakeIPv6Route}) + cr = append(cr, m.enableFakeIPRoutes()...) } m.notifier.SetInitialClientRoutes(cr, routesForComparison) } +func (m *DefaultManager) enableFakeIPRoutes() []*route.Route { + m.fakeIPManager = fakeip.NewManager() + + v4ID := uuid.NewString() + fakeIPRoute := &route.Route{ + ID: route.ID(v4ID), + Network: m.fakeIPManager.GetFakeIPBlock(), + NetID: route.NetID(v4ID), + Peer: m.pubKey, + NetworkType: route.IPv4Network, + } + v6ID := uuid.NewString() + fakeIPv6Route := &route.Route{ + ID: route.ID(v6ID), + Network: m.fakeIPManager.GetFakeIPv6Block(), + NetID: route.NetID(v6ID), + Peer: m.pubKey, + NetworkType: route.IPv6Network, + } + fakeRoutes := []*route.Route{fakeIPRoute, fakeIPv6Route} + m.notifier.SetFakeIPRoutes(fakeRoutes) + return fakeRoutes +} + func (m *DefaultManager) setupRefCounters(useNoop bool) { var once sync.Once var wgIface *net.Interface @@ -464,6 +469,9 @@ func (m *DefaultManager) UpdateRoutes( var merr *multierror.Error if !m.disableClientRoutes { + if runtime.GOOS == "android" && useNewDNSRoute && m.fakeIPManager == nil { + m.enableFakeIPRoutes() + } // Update route selector based on management server's isSelected status m.updateRouteSelectorFromManagement(clientRoutes) diff --git a/client/internal/routemanager/notifier/notifier_android.go b/client/internal/routemanager/notifier/notifier_android.go index 49300dbb2..60e1d0a0f 100644 --- a/client/internal/routemanager/notifier/notifier_android.go +++ b/client/internal/routemanager/notifier/notifier_android.go @@ -41,6 +41,7 @@ func (n *Notifier) SetInitialClientRoutes(initialRoutes []*route.Route, routesFo // SetFakeIPRoutes stores the fake IP routes to be included in every TUN rebuild. func (n *Notifier) SetFakeIPRoutes(routes []*route.Route) { n.fakeIPRoutes = routes + n.notify() } func (n *Notifier) OnNewRoutes(idMap route.HAMap) { @@ -78,9 +79,7 @@ func (n *Notifier) notify() { routeStrings := n.routesToStrings(allRoutes) sort.Strings(routeStrings) - go func(l listener.NetworkChangeListener) { - l.OnNetworkChanged(strings.Join(routeStrings, ",")) - }(n.listener) + n.listener.OnNetworkChanged(strings.Join(routeStrings, ",")) } func filterStatic(routes []*route.Route) []*route.Route { @@ -102,16 +101,11 @@ func (n *Notifier) routesToStrings(routes []*route.Route) []string { } func (n *Notifier) hasRouteDiff(a []*route.Route, b []*route.Route) bool { - slices.SortFunc(a, func(x, y *route.Route) int { - return strings.Compare(x.NetString(), y.NetString()) - }) - slices.SortFunc(b, func(x, y *route.Route) int { - return strings.Compare(x.NetString(), y.NetString()) - }) - - return !slices.EqualFunc(a, b, func(x, y *route.Route) bool { - return x.NetString() == y.NetString() - }) + as := n.routesToStrings(a) + bs := n.routesToStrings(b) + sort.Strings(as) + sort.Strings(bs) + return !slices.Equal(as, bs) } func (n *Notifier) GetInitialRouteRanges() []string { diff --git a/client/internal/updater/installer/installer_run_darwin.go b/client/internal/updater/installer/installer_run_darwin.go index 248a404aa..5650bc769 100644 --- a/client/internal/updater/installer/installer_run_darwin.go +++ b/client/internal/updater/installer/installer_run_darwin.go @@ -98,47 +98,44 @@ func (u *Installer) startDaemon(daemonFolder string) error { func (u *Installer) startUIAsUser() error { log.Infof("starting netbird-ui: %s", uiBinary) - // Get the current console user - cmd := exec.Command("stat", "-f", "%Su", "/dev/console") - output, err := cmd.Output() + username, err := consoleUser() if err != nil { - return fmt.Errorf("failed to get console user: %w", err) + return err } - username := strings.TrimSpace(string(output)) - if username == "" || username == "root" { - return fmt.Errorf("no active user session found") - } - - log.Infof("starting UI for user: %s", username) - - // Get user's UID userInfo, err := user.Lookup(username) if err != nil { - return fmt.Errorf("failed to lookup user %s: %w", username, err) + return fmt.Errorf("lookup user %s: %w", username, err) } - // Start the UI process as the console user using launchctl - // This ensures the app runs in the user's context with proper GUI access - launchCmd := exec.Command("launchctl", "asuser", userInfo.Uid, "open", "-a", uiBinary) + log.Infof("starting UI for user: %s (uid %s)", username, userInfo.Uid) + + launchCmd := exec.Command("launchctl", "asuser", userInfo.Uid, "sudo", "-u", username, "-H", "open", "-a", uiBinary) log.Infof("launchCmd: %s", launchCmd.String()) - // Set the user's home directory for proper macOS app behavior - launchCmd.Env = append(os.Environ(), "HOME="+userInfo.HomeDir) - log.Infof("set HOME environment variable: %s", userInfo.HomeDir) - if err := launchCmd.Start(); err != nil { - return fmt.Errorf("failed to start UI process: %w", err) - } - - // Release the process so it can run independently - if err := launchCmd.Process.Release(); err != nil { - log.Warnf("failed to release UI process: %v", err) + if err := launchCmd.Run(); err != nil { + return fmt.Errorf("run UI launch: %w", err) } log.Infof("netbird-ui started successfully for user %s", username) return nil } +func consoleUser() (string, error) { + output, err := exec.Command("stat", "-f", "%Su", "/dev/console").Output() + if err != nil { + return "", fmt.Errorf("get console user: %w", err) + } + + username := strings.TrimSpace(string(output)) + switch username { + case "", "root", "loginwindow", "_mbsetupuser": + return "", fmt.Errorf("no active GUI user session, console user: %q", username) + } + + return username, nil +} + func (u *Installer) installPkgFile(ctx context.Context, path string) error { log.Infof("installing pkg file: %s", path) diff --git a/client/ios/NetBirdSDK/client.go b/client/ios/NetBirdSDK/client.go index 37d3e5d99..2d5460d03 100644 --- a/client/ios/NetBirdSDK/client.go +++ b/client/ios/NetBirdSDK/client.go @@ -158,13 +158,19 @@ func (c *Client) Run(fd int32, interfaceName string, envList *EnvList) error { defer c.ctxCancel() c.ctxCancelLock.Unlock() - auth := NewAuthWithConfig(ctx, cfg) - err = auth.LoginSync() - if err != nil { - return err - } - - log.Infof("Auth successful") + // No login pre-flight here. The engine's own loginToManagement (connect.go) performs + // the authoritative Login immediately before the first Sync, so a LoginSync() call at + // this point only duplicated it — costing two extra Login RPCs (IsLoginRequired + + // Login) on every engine start, since IsLoginRequired is itself a full Login RPC. + // + // Auth failures still reach the caller through the engine path: loginToManagement + // returns PermissionDenied, which marks the shared status recorder + // (MarkManagementDisconnected) and fires ClientStop → onDisconnected, where + // IsLoginRequiredCached() reports login-required. The error is also returned out of Run(). + // + // A pre-flight was also actively harmful when the server is unreachable: its 2-minute + // backoff blocked the start and then reported "login required" for what was really a + // timeout. The engine instead keeps retrying and recovers when the server returns. // todo do not throw error in case of cancelled context ctx = internal.CtxInitState(ctx) c.onHostDnsFn = func([]string) {} diff --git a/client/ios/NetBirdSDK/login.go b/client/ios/NetBirdSDK/login.go index 99486839b..6cba0c411 100644 --- a/client/ios/NetBirdSDK/login.go +++ b/client/ios/NetBirdSDK/login.go @@ -222,17 +222,36 @@ func (a *Auth) Login(resultListener ErrListener, urlOpener URLOpener, forceDevic // LoginWithDeviceName performs interactive login with device authentication support // The deviceName parameter allows specifying a custom device name (required for tvOS) func (a *Auth) LoginWithDeviceName(resultListener ErrListener, urlOpener URLOpener, forceDeviceAuth bool, deviceName string) { + a.startLogin(resultListener, urlOpener, forceDeviceAuth, deviceName, false) +} + +// LoginInteractive performs the same interactive login as LoginWithDeviceName but skips the +// IsLoginRequired() pre-flight and goes straight to the browser / device-code flow. +// +// IsLoginRequired() is itself a full Login RPC against the management server, so when the +// caller has ALREADY established that login is required it is a pure duplicate. On iOS the +// main app decides to show the browser based on its own isLoginRequired() check and then +// calls straight into this method, so re-asking the server would add another Login RPC to +// every interactive login. +// +// Use LoginWithDeviceName when the auth state is unknown and a silent (browser-less) login +// must still be possible; use this when the browser is going to be shown regardless. +func (a *Auth) LoginInteractive(resultListener ErrListener, urlOpener URLOpener, forceDeviceAuth bool, deviceName string) { + a.startLogin(resultListener, urlOpener, forceDeviceAuth, deviceName, true) +} + +func (a *Auth) startLogin(resultListener ErrListener, urlOpener URLOpener, forceDeviceAuth bool, deviceName string, skipLoginCheck bool) { if resultListener == nil { - log.Errorf("LoginWithDeviceName: resultListener is nil") + log.Errorf("startLogin: resultListener is nil") return } if urlOpener == nil { - log.Errorf("LoginWithDeviceName: urlOpener is nil") + log.Errorf("startLogin: urlOpener is nil") resultListener.OnError(fmt.Errorf("urlOpener is nil")) return } go func() { - err := a.login(urlOpener, forceDeviceAuth, deviceName) + err := a.login(urlOpener, forceDeviceAuth, deviceName, skipLoginCheck) if err != nil { resultListener.OnError(err) } else { @@ -241,7 +260,7 @@ func (a *Auth) LoginWithDeviceName(resultListener ErrListener, urlOpener URLOpen }() } -func (a *Auth) login(urlOpener URLOpener, forceDeviceAuth bool, deviceName string) error { +func (a *Auth) login(urlOpener URLOpener, forceDeviceAuth bool, deviceName string, skipLoginCheck bool) error { // Create context with device name if provided ctx := a.ctx if deviceName != "" { @@ -255,10 +274,13 @@ func (a *Auth) login(urlOpener URLOpener, forceDeviceAuth bool, deviceName strin } defer authClient.Close() - // check if we need to generate JWT token - needsLogin, err := authClient.IsLoginRequired(ctx) - if err != nil { - return fmt.Errorf("failed to check login requirement: %v", err) + // check if we need to generate JWT token (skipped when the caller already knows) + needsLogin := true + if !skipLoginCheck { + needsLogin, err = authClient.IsLoginRequired(ctx) + if err != nil { + return fmt.Errorf("failed to check login requirement: %v", err) + } } jwtToken := "" diff --git a/client/server/login_outcome_test.go b/client/server/login_outcome_test.go new file mode 100644 index 000000000..d3b1b7c52 --- /dev/null +++ b/client/server/login_outcome_test.go @@ -0,0 +1,89 @@ +package server + +import ( + "context" + "encoding/json" + "errors" + "os" + "testing" + + "github.com/stretchr/testify/require" + "google.golang.org/grpc/codes" + gstatus "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/internal" + "github.com/netbirdio/netbird/client/proto" +) + +// A login that never reached Management is not a decision about the peer's +// credentials, so it must come back as a retryable error rather than an SSO +// prompt: the user cannot finish a browser login while Management is down, and +// the CLI's own backoff resolves the outage on its own once the daemon reports +// the failure. Reproduces `netbird down; netbird up` printing a device-code URL +// because Management happened to be restarting when the daemon dialed it. +func TestLogin_ManagementUnreachableIsReturnedInsteadOfDemandingSSO(t *testing.T) { + s, _, _, username, _ := setupServerWithProfile(t) + s.rootCtx = internal.CtxInitState(context.Background()) + + unreachable := errors.New("create connection: dial context: context deadline exceeded") + attempts := 0 + s.loginAttemptFn = func(context.Context, string, string) (internal.StatusType, error) { + attempts++ + return internal.StatusLoginFailed, unreachable + } + + resp, err := s.Login(userCtx(), &proto.LoginRequest{Username: &username}) + require.Error(t, err) + require.ErrorIs(t, err, unreachable, "the transport failure was replaced by something else") + require.Nil(t, resp, "a failed login must not answer with a login response") + require.Equal(t, 1, attempts) + require.Nil(t, s.oauthAuthFlow.flow, "the daemon started an SSO flow for a peer whose login was never decided") + + status, err := internal.CtxGetState(s.rootCtx).Status() + require.NoError(t, err) + require.Equal(t, internal.StatusLoginFailed, status, + "a peer that could not reach Management is not waiting on a login") +} + +// The counterpart: Management refusing the peer's credentials is a decision, and +// the SSO flow still has to start for it. The profile carries an unusable +// private key so the flow setup fails immediately instead of dialing, which is +// enough to show the branch was entered — the refusal itself is never what comes +// back out. +func TestLogin_AuthRefusalStartsSSOFlow(t *testing.T) { + s, _, _, username, cfgPath := setupServerWithProfile(t) + s.rootCtx = internal.CtxInitState(context.Background()) + breakProfilePrivateKey(t, cfgPath) + + refused := gstatus.Error(codes.PermissionDenied, "peer is not registered") + s.loginAttemptFn = func(context.Context, string, string) (internal.StatusType, error) { + return internal.StatusNeedsLogin, refused + } + + _, err := s.Login(userCtx(), &proto.LoginRequest{Username: &username}) + require.Error(t, err) + require.NotErrorIs(t, err, refused, + "the refusal was handed back to the caller instead of starting the SSO flow") + + status, stateErr := internal.CtxGetState(s.rootCtx).Status() + require.NoError(t, stateErr) + require.Equal(t, internal.StatusLoginFailed, status, + "the SSO flow setup was never reached with the broken key") +} + +// breakProfilePrivateKey replaces the profile's private key with an unparseable +// one, which makes any attempt to build a Management client fail on the spot. +func breakProfilePrivateKey(t *testing.T, cfgPath string) { + t.Helper() + + raw, err := os.ReadFile(cfgPath) + require.NoError(t, err) + + var cfg map[string]any + require.NoError(t, json.Unmarshal(raw, &cfg)) + cfg["PrivateKey"] = "not-a-key" + + patched, err := json.Marshal(cfg) + require.NoError(t, err) + require.NoError(t, os.WriteFile(cfgPath, patched, 0o600)) +} diff --git a/client/server/server.go b/client/server/server.go index aaab5cc02..892c9c5de 100644 --- a/client/server/server.go +++ b/client/server/server.go @@ -135,6 +135,11 @@ type Server struct { updateManager *updater.Manager jwtCache *jwtCache + + // loginAttemptFn stands in for the Management login round trip. Tests set + // it to drive the login outcomes that need a server on the other end; + // production leaves it nil, and every login goes through loginAttempt. + loginAttemptFn func(ctx context.Context, setupKey, jwtToken string) (internal.StatusType, error) } type oauthAuthFlow struct { @@ -370,7 +375,19 @@ func (s *Server) connectionGoroutineRunning() bool { } } -// loginAttempt attempts to login using the provided information. it returns a status in case something fails +// attemptLogin runs a login round trip against Management, or the stand-in a +// test installed in place of it. +func (s *Server) attemptLogin(ctx context.Context, setupKey, jwtToken string) (internal.StatusType, error) { + if s.loginAttemptFn != nil { + return s.loginAttemptFn(ctx, setupKey, jwtToken) + } + return s.loginAttempt(ctx, setupKey, jwtToken) +} + +// loginAttempt attempts to login using the provided information. It returns +// StatusNeedsLogin when Management refused the peer's credentials and +// StatusLoginFailed for every other failure, so callers can tell an +// authentication decision apart from a login that never got made. func (s *Server) loginAttempt(ctx context.Context, setupKey, jwtToken string) (internal.StatusType, error) { authClient, err := auth.NewAuth(ctx, s.config.PrivateKey, s.config.ManagementURL, s.config) if err != nil { @@ -623,11 +640,23 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro s.config = config s.mutex.Unlock() - if _, err := s.loginAttempt(ctx, "", ""); err == nil { + loginStatus, err := s.attemptLogin(ctx, "", "") + if err == nil { state.Set(internal.StatusIdle) return &proto.LoginResponse{}, nil } + // Only an authentication refusal means the peer has to (re-)authenticate. + // Any other failure leaves the login undecided: Management unreachable, a + // restart mid-request, an internal error. Those are returned for the caller + // to retry, because turning them into an SSO prompt asks the user to solve + // something that is not theirs to solve, and a browser login cannot succeed + // while Management is unreachable anyway. + if loginStatus != internal.StatusNeedsLogin { + state.Set(loginStatus) + return nil, err + } + if msg.SetupKey == "" { hint := "" if msg.Hint != nil { @@ -684,7 +713,7 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro // which returns NeedsLogin and parks on the browser leg. state.Set(internal.StatusConnecting) - if loginStatus, err := s.loginAttempt(ctx, msg.SetupKey, ""); err != nil { + if loginStatus, err := s.attemptLogin(ctx, msg.SetupKey, ""); err != nil { state.Set(loginStatus) return nil, err } @@ -839,7 +868,7 @@ func (s *Server) WaitSSOLogin(callerCtx context.Context, msg *proto.WaitSSOLogin s.oauthAuthFlow.expiresAt = time.Now() s.mutex.Unlock() - if loginStatus, err := s.loginAttempt(ctx, "", tokenInfo.GetTokenToUse()); err != nil { + if loginStatus, err := s.attemptLogin(ctx, "", tokenInfo.GetTokenToUse()); err != nil { state.Set(loginStatus) return nil, err } diff --git a/client/ui/build/linux/nfpm/nfpm.yaml b/client/ui/build/linux/nfpm/nfpm.yaml index a05daef62..764855a63 100644 --- a/client/ui/build/linux/nfpm/nfpm.yaml +++ b/client/ui/build/linux/nfpm/nfpm.yaml @@ -26,17 +26,17 @@ contents: # Default dependencies for the GTK4 + WebKitGTK 6.0 stack (Ubuntu 24.04+ / Debian 13+) depends: - - libgtk-4-1 + - libgtk-4-1 (>= 4.14) - libwebkitgtk-6.0-4 - xdg-utils # Distribution-specific overrides for different package formats overrides: - # RPM packages for Fedora / RHEL / AlmaLinux / Rocky Linux + # RPM packages for Fedora / RHEL / AlmaLinux / Rocky Linux / openSUSE rpm: depends: - - gtk4 - - webkitgtk6.0 + - (gtk4 >= 4.14 or libgtk-4-1 >= 4.14) + - (webkitgtk6.0 or libwebkitgtk-6_0-4) - xdg-utils # Arch Linux packages diff --git a/client/ui/frontend/src/lib/connection.ts b/client/ui/frontend/src/lib/connection.ts index fca03fc87..cc0e67cb3 100644 --- a/client/ui/frontend/src/lib/connection.ts +++ b/client/ui/frontend/src/lib/connection.ts @@ -43,7 +43,12 @@ function buildSsoCancelPromise(state: SsoState, signal?: AbortSignal): Promise { @@ -56,7 +61,7 @@ async function runSsoLogin( // suspended, so a frontend-driven Up (a promise continuation) would not // fire until the user woke the window (e.g. hovering the tray icon). const waitPromise = Connection.WaitSSOLoginAndUp( - { userCode: result.userCode, hostname: "" }, + { userCode: result.userCode, hostname: "", profileId: result.profileId }, { profileName: "", username: "" }, ); diff --git a/client/ui/main.go b/client/ui/main.go index 782e76d1c..e2d172e5b 100644 --- a/client/ui/main.go +++ b/client/ui/main.go @@ -14,7 +14,6 @@ import ( "github.com/sirupsen/logrus" "github.com/wailsapp/wails/v3/pkg/application" "github.com/wailsapp/wails/v3/pkg/events" - "github.com/wailsapp/wails/v3/pkg/services/notifications" "github.com/netbirdio/netbird/client/ui/authsession" "github.com/netbirdio/netbird/client/ui/i18n" @@ -63,7 +62,7 @@ type registeredServices struct { profiles *services.Profiles update *services.Update daemonFeed *services.DaemonFeed - notifier *notifications.NotificationService + notifier *Notifier compat *services.Compat profileSwitcher *services.ProfileSwitcher bundle *i18n.Bundle @@ -102,7 +101,7 @@ func main() { updaterHolder := updater.NewHolder(app.Event) update := services.NewUpdate(conn, updaterHolder) daemonFeed := services.NewDaemonFeed(conn, app.Event, updaterHolder, debugLog) - notifier := notifications.New() + notifier := newNotifier() compat := services.NewCompat(conn) // macOS shows no toast until permission is requested. Run it after // ApplicationStarted so the notifier's Startup has initialised the @@ -210,7 +209,7 @@ func main() { // requestNotificationAuthorization prompts for macOS notification permission. // The request blocks until the user responds (up to 3 minutes), so callers run // it in a goroutine. No-op on Linux/Windows. -func requestNotificationAuthorization(notifier *notifications.NotificationService) { +func requestNotificationAuthorization(notifier *Notifier) { authorized, err := notifier.CheckNotificationAuthorization() if err != nil { logrus.Debugf("check notification authorization: %v", err) diff --git a/client/ui/notifier.go b/client/ui/notifier.go new file mode 100644 index 000000000..71ae3b0df --- /dev/null +++ b/client/ui/notifier.go @@ -0,0 +1,101 @@ +//go:build !android && !ios && !freebsd && !js + +package main + +import ( + "context" + "errors" + "sync/atomic" + + log "github.com/sirupsen/logrus" + "github.com/wailsapp/wails/v3/pkg/application" + "github.com/wailsapp/wails/v3/pkg/services/notifications" +) + +var errNotificationsUnavailable = errors.New("notifications unavailable") + +// Notifier wraps the Wails notification service so an unavailable backend +// disables notifications instead of aborting the app. Startup fails for +// environment reasons (a bare unbundled binary on macOS has no bundle +// identifier, a headless Linux session has no D-Bus session bus), and Wails +// treats a service startup error as fatal. After a failed startup every call +// is a no-op: on macOS, touching UNUserNotificationCenter without a bundle +// identifier raises an Objective-C exception that recover() cannot catch. +type Notifier struct { + inner *notifications.NotificationService + available atomic.Bool +} + +func newNotifier() *Notifier { + return &Notifier{inner: notifications.New()} +} + +// ServiceName implements the Wails service-name hook for startup logs. +func (n *Notifier) ServiceName() string { + return n.inner.ServiceName() +} + +// ServiceStartup starts the platform notifier, downgrading failure to a +// warning so the app keeps running without notifications. +func (n *Notifier) ServiceStartup(ctx context.Context, options application.ServiceOptions) error { + if err := n.inner.ServiceStartup(ctx, options); err != nil { + log.Warnf("notifications disabled: %v", err) + return nil + } + n.available.Store(true) + return nil +} + +func (n *Notifier) ServiceShutdown() error { + if !n.available.Load() { + return nil + } + return n.inner.ServiceShutdown() +} + +func (n *Notifier) CheckNotificationAuthorization() (bool, error) { + if !n.available.Load() { + return false, errNotificationsUnavailable + } + return n.inner.CheckNotificationAuthorization() +} + +func (n *Notifier) RequestNotificationAuthorization() (bool, error) { + if !n.available.Load() { + return false, errNotificationsUnavailable + } + return n.inner.RequestNotificationAuthorization() +} + +// SendNotification delivers a notification, silently dropping it when the +// backend never started (notifications are best-effort everywhere). +func (n *Notifier) SendNotification(options notifications.NotificationOptions) error { + if !n.available.Load() { + log.Debugf("notifications disabled, dropping %q", options.ID) + return nil + } + return n.inner.SendNotification(options) +} + +func (n *Notifier) SendNotificationWithActions(options notifications.NotificationOptions) error { + if !n.available.Load() { + log.Debugf("notifications disabled, dropping %q", options.ID) + return nil + } + return n.inner.SendNotificationWithActions(options) +} + +func (n *Notifier) RegisterNotificationCategory(category notifications.NotificationCategory) error { + if !n.available.Load() { + return nil + } + return n.inner.RegisterNotificationCategory(category) +} + +// OnNotificationResponse registers the response callback. Pure Go state, so +// it is safe (and simply inert) when the backend never started. +// +//wails:ignore +func (n *Notifier) OnNotificationResponse(callback func(result notifications.NotificationResult)) { + n.inner.OnNotificationResponse(callback) +} diff --git a/client/ui/preferences/store.go b/client/ui/preferences/store.go index df6fbbb16..49acb7917 100644 --- a/client/ui/preferences/store.go +++ b/client/ui/preferences/store.go @@ -246,6 +246,7 @@ func (s *Store) ExistedAtLoad() bool { func (s *Store) load() error { if _, err := os.Stat(s.path); err != nil { if errors.Is(err, os.ErrNotExist) { + log.Infof("no ui preferences file at %s; using defaults", s.path) return nil } return fmt.Errorf("stat preferences: %w", err) diff --git a/client/ui/services/connection.go b/client/ui/services/connection.go index fae7ddd23..1069f8754 100644 --- a/client/ui/services/connection.go +++ b/client/ui/services/connection.go @@ -33,12 +33,21 @@ type LoginResult struct { UserCode string `json:"userCode"` VerificationURI string `json:"verificationUri"` VerificationURIComplete string `json:"verificationUriComplete"` + // ProfileID is the ID of the profile this login ran against, or "" when the + // caller named the profile itself and no ID was resolved. Pass it back in + // WaitSSOParams so the account email lands on this profile even if the + // active one changes during SSO. + ProfileID string `json:"profileId"` } // WaitSSOParams are the inputs to waitSSOLogin. type WaitSSOParams struct { UserCode string `json:"userCode"` Hostname string `json:"hostname"` + // ProfileID is the profile the login was started for, used to file the + // account email against it rather than against whichever profile is active + // when the flow returns. Optional: empty falls back to the active profile. + ProfileID string `json:"profileId"` } // UpParams selects the profile to bring up. @@ -77,11 +86,16 @@ func (s *Connection) Login(ctx context.Context, p LoginParams) (LoginResult, err // Fall back to the daemon's active profile and the current OS user. profileName := p.ProfileName username := p.Username + // Only set when the daemon told us the ID. A caller-supplied ProfileName is + // a handle — a display name or an ID prefix resolve too — and the state file + // is named after the ID, so passing a handle on would name the wrong file. + profileID := "" if profileName == "" { if active, aerr := cli.GetActiveProfile(ctx, &proto.GetActiveProfileRequest{}); aerr == nil { // Address the active profile by ID (the daemon resolves it as a // handle); names can collide, the ID cannot. profileName = active.GetId() + profileID = profileName if username == "" { username = active.GetUsername() } @@ -122,6 +136,7 @@ func (s *Connection) Login(ctx context.Context, p LoginParams) (LoginResult, err UserCode: resp.GetUserCode(), VerificationURI: resp.GetVerificationURI(), VerificationURIComplete: resp.GetVerificationURIComplete(), + ProfileID: profileID, }, nil } @@ -242,6 +257,31 @@ func (s *Connection) waitSSOLogin(ctx context.Context, p WaitSSOParams) (string, return "", s.classifyDaemonError(err) } log.Infof("SSO login completed, daemon reported success") + + // Persist the account email the same way the CLI does after its own + // WaitSSOLogin: the daemon returns it but cannot store it, since it runs as + // root and the per-profile state file is user-owned (see Logout below). + // Without this the profile has no email, so Profiles.List shows no account + // and later logins and session extends go out without a login_hint — + // leaving the IdP to guess which account was meant. + if email := resp.GetEmail(); email != "" { + state := &profilemanager.ProfileState{Email: email} + pm := profilemanager.NewProfileManager() + + // Against the profile the login was started for: SSO spans seconds of + // user interaction, and a profile switch in that window would otherwise + // file the email under the wrong profile. + if p.ProfileID != "" { + err = pm.SetProfileState(profilemanager.ID(p.ProfileID), state) + } else { + err = pm.SetActiveProfileState(state) + } + if err != nil { + // Non-fatal: the login itself succeeded. + log.Warnf("failed to store account email: %v", err) + } + } + return resp.GetEmail(), nil } diff --git a/client/ui/services/profile.go b/client/ui/services/profile.go index 09468c9df..5a9a0e68d 100644 --- a/client/ui/services/profile.go +++ b/client/ui/services/profile.go @@ -6,6 +6,8 @@ import ( "context" "os/user" + log "github.com/sirupsen/logrus" + "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/proto" ) @@ -151,11 +153,31 @@ func (s *Profiles) Remove(ctx context.Context, p ProfileRef) error { if err != nil { return err } - _, err = cli.RemoveProfile(ctx, &proto.RemoveProfileRequest{ + resp, err := cli.RemoveProfile(ctx, &proto.RemoveProfileRequest{ ProfileName: p.ProfileName, Username: p.Username, }) - return err + if err != nil { + return err + } + + // The daemon deletes what it owns but runs as root, so it leaves the + // user-owned state file holding the account email behind (same split as + // Connection.Logout). Legacy profiles are keyed by name rather than by a + // generated ID, so a recreated profile of the same name would inherit the + // deleted one's email and offer it as the login_hint. + // + // Keyed on the ID the daemon resolved, not on the request handle: that may + // have been a display name or an ID prefix, which would name a different + // file (or none). + if id := resp.GetId(); id != "" { + if err := profilemanager.NewProfileManager().RemoveProfileState(id); err != nil { + // Non-fatal: the profile itself is gone. + log.Warnf("failed to remove profile state for %s: %v", id, err) + } + } + + return nil } // Rename changes a profile's display name. The on-disk ID is unaffected, so diff --git a/client/ui/tray.go b/client/ui/tray.go index 3050d159a..3093c693b 100644 --- a/client/ui/tray.go +++ b/client/ui/tray.go @@ -44,7 +44,7 @@ type TrayServices struct { Profiles *services.Profiles Networks *services.Networks DaemonFeed *services.DaemonFeed - Notifier *notifications.NotificationService + Notifier *Notifier Update *services.Update ProfileSwitcher *services.ProfileSwitcher WindowManager *services.WindowManager diff --git a/client/ui/tray_notify.go b/client/ui/tray_notify.go index d1117b57b..5b2629419 100644 --- a/client/ui/tray_notify.go +++ b/client/ui/tray_notify.go @@ -44,7 +44,7 @@ func safeSendNotification(send sendFn, what string, opts notifications.Notificat // notifyIfDaemonOutdated probes the daemon once and fires an OS toast when it // is reachable but too old for this UI. A probe error means the daemon isn't // reachable (not outdated), so it is left to the normal connection flow. -func notifyIfDaemonOutdated(compat *services.Compat, notifier *notifications.NotificationService, loc *Localizer) { +func notifyIfDaemonOutdated(compat *services.Compat, notifier *Notifier, loc *Localizer) { ready, err := compat.DaemonReady(context.Background()) if err != nil { log.Debugf("daemon compatibility probe: %v", err) diff --git a/client/ui/tray_update.go b/client/ui/tray_update.go index 1a377dfa3..27037eccb 100644 --- a/client/ui/tray_update.go +++ b/client/ui/tray_update.go @@ -21,7 +21,7 @@ type trayUpdater struct { app *application.App window *application.WebviewWindow update *services.Update - notifier *notifications.NotificationService + notifier *Notifier loc *Localizer onIconChange func() // onMenuChange drives a full tray relayout: the update row lives in the @@ -36,7 +36,7 @@ type trayUpdater struct { progressWindowOpen bool } -func newTrayUpdater(app *application.App, window *application.WebviewWindow, update *services.Update, notifier *notifications.NotificationService, loc *Localizer, onIconChange func(), onMenuChange func()) *trayUpdater { +func newTrayUpdater(app *application.App, window *application.WebviewWindow, update *services.Update, notifier *Notifier, loc *Localizer, onIconChange func(), onMenuChange func()) *trayUpdater { u := &trayUpdater{ app: app, window: window, diff --git a/management/internals/modules/agentnetwork/manager.go b/management/internals/modules/agentnetwork/manager.go index 2e39bde5d..ba2c06826 100644 --- a/management/internals/modules/agentnetwork/manager.go +++ b/management/internals/modules/agentnetwork/manager.go @@ -157,14 +157,14 @@ func NewManager( } func (m *managerImpl) GetAllProviders(ctx context.Context, accountID, userID string) ([]*types.Provider, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Read); err != nil { return nil, err } return m.store.GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, accountID) } func (m *managerImpl) GetProvider(ctx context.Context, accountID, userID, providerID string) (*types.Provider, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Read); err != nil { return nil, err } return m.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, accountID, providerID) @@ -175,9 +175,14 @@ func (m *managerImpl) GetProvider(ctx context.Context, accountID, userID, provid // been created yet; otherwise it is ignored (the cluster is pinned on // Settings and every provider in the account routes through it). func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provider *types.Provider, bootstrapCluster string) (*types.Provider, error) { - if err := m.requirePermission(ctx, provider.AccountID, userID, operations.Create); err != nil { + if err := m.requirePermission(ctx, provider.AccountID, userID, modules.AgentNetworkProviders, operations.Create); err != nil { return nil, err } + if strings.TrimSpace(bootstrapCluster) != "" { + if err := m.requireSettingsBootstrapPermission(ctx, provider.AccountID, userID); err != nil { + return nil, err + } + } // An empty api_key would silently produce a synthesised service // that 401s on every upstream request. Surface the misconfiguration @@ -218,7 +223,7 @@ func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provide } func (m *managerImpl) UpdateProvider(ctx context.Context, userID string, provider *types.Provider) (*types.Provider, error) { - if err := m.requirePermission(ctx, provider.AccountID, userID, operations.Update); err != nil { + if err := m.requirePermission(ctx, provider.AccountID, userID, modules.AgentNetworkProviders, operations.Update); err != nil { return nil, err } @@ -257,7 +262,7 @@ func (m *managerImpl) UpdateProvider(ctx context.Context, userID string, provide } func (m *managerImpl) DeleteProvider(ctx context.Context, accountID, userID, providerID string) error { - if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Delete); err != nil { return err } @@ -306,21 +311,21 @@ func pluralize(n int, singular, plural string) string { } func (m *managerImpl) GetAllPolicies(ctx context.Context, accountID, userID string) ([]*types.Policy, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkPolicies, operations.Read); err != nil { return nil, err } return m.store.GetAccountAgentNetworkPolicies(ctx, store.LockingStrengthNone, accountID) } func (m *managerImpl) GetPolicy(ctx context.Context, accountID, userID, policyID string) (*types.Policy, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkPolicies, operations.Read); err != nil { return nil, err } return m.store.GetAgentNetworkPolicyByID(ctx, store.LockingStrengthNone, accountID, policyID) } func (m *managerImpl) CreatePolicy(ctx context.Context, userID string, policy *types.Policy) (*types.Policy, error) { - if err := m.requirePermission(ctx, policy.AccountID, userID, operations.Create); err != nil { + if err := m.requirePermission(ctx, policy.AccountID, userID, modules.AgentNetworkPolicies, operations.Create); err != nil { return nil, err } @@ -346,7 +351,7 @@ func (m *managerImpl) CreatePolicy(ctx context.Context, userID string, policy *t } func (m *managerImpl) UpdatePolicy(ctx context.Context, userID string, policy *types.Policy) (*types.Policy, error) { - if err := m.requirePermission(ctx, policy.AccountID, userID, operations.Update); err != nil { + if err := m.requirePermission(ctx, policy.AccountID, userID, modules.AgentNetworkPolicies, operations.Update); err != nil { return nil, err } @@ -373,7 +378,7 @@ func (m *managerImpl) UpdatePolicy(ctx context.Context, userID string, policy *t } func (m *managerImpl) DeletePolicy(ctx context.Context, accountID, userID, policyID string) error { - if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkPolicies, operations.Delete); err != nil { return err } @@ -393,21 +398,21 @@ func (m *managerImpl) DeletePolicy(ctx context.Context, accountID, userID, polic } func (m *managerImpl) GetAllGuardrails(ctx context.Context, accountID, userID string) ([]*types.Guardrail, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkGuardrails, operations.Read); err != nil { return nil, err } return m.store.GetAccountAgentNetworkGuardrails(ctx, store.LockingStrengthNone, accountID) } func (m *managerImpl) GetGuardrail(ctx context.Context, accountID, userID, guardrailID string) (*types.Guardrail, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkGuardrails, operations.Read); err != nil { return nil, err } return m.store.GetAgentNetworkGuardrailByID(ctx, store.LockingStrengthNone, accountID, guardrailID) } func (m *managerImpl) CreateGuardrail(ctx context.Context, userID string, guardrail *types.Guardrail) (*types.Guardrail, error) { - if err := m.requirePermission(ctx, guardrail.AccountID, userID, operations.Create); err != nil { + if err := m.requirePermission(ctx, guardrail.AccountID, userID, modules.AgentNetworkGuardrails, operations.Create); err != nil { return nil, err } @@ -429,7 +434,7 @@ func (m *managerImpl) CreateGuardrail(ctx context.Context, userID string, guardr } func (m *managerImpl) UpdateGuardrail(ctx context.Context, userID string, guardrail *types.Guardrail) (*types.Guardrail, error) { - if err := m.requirePermission(ctx, guardrail.AccountID, userID, operations.Update); err != nil { + if err := m.requirePermission(ctx, guardrail.AccountID, userID, modules.AgentNetworkGuardrails, operations.Update); err != nil { return nil, err } @@ -452,7 +457,7 @@ func (m *managerImpl) UpdateGuardrail(ctx context.Context, userID string, guardr } func (m *managerImpl) DeleteGuardrail(ctx context.Context, accountID, userID, guardrailID string) error { - if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkGuardrails, operations.Delete); err != nil { return err } @@ -473,7 +478,7 @@ func (m *managerImpl) DeleteGuardrail(ctx context.Context, accountID, userID, gu // GetAllBudgetRules returns every account-level budget rule for the account. func (m *managerImpl) GetAllBudgetRules(ctx context.Context, accountID, userID string) ([]*types.AccountBudgetRule, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkBudgets, operations.Read); err != nil { return nil, err } return m.store.GetAccountAgentNetworkBudgetRules(ctx, store.LockingStrengthNone, accountID) @@ -481,7 +486,7 @@ func (m *managerImpl) GetAllBudgetRules(ctx context.Context, accountID, userID s // GetBudgetRule returns a single account-level budget rule. func (m *managerImpl) GetBudgetRule(ctx context.Context, accountID, userID, ruleID string) (*types.AccountBudgetRule, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkBudgets, operations.Read); err != nil { return nil, err } return m.store.GetAgentNetworkBudgetRuleByID(ctx, store.LockingStrengthNone, accountID, ruleID) @@ -491,7 +496,7 @@ func (m *managerImpl) GetBudgetRule(ctx context.Context, accountID, userID, rule // enforced at request time (CheckLLMPolicyLimits), not baked into the synth // proxy config, so no reconcile is needed. func (m *managerImpl) CreateBudgetRule(ctx context.Context, userID string, rule *types.AccountBudgetRule) (*types.AccountBudgetRule, error) { - if err := m.requirePermission(ctx, rule.AccountID, userID, operations.Create); err != nil { + if err := m.requirePermission(ctx, rule.AccountID, userID, modules.AgentNetworkBudgets, operations.Create); err != nil { return nil, err } @@ -513,7 +518,7 @@ func (m *managerImpl) CreateBudgetRule(ctx context.Context, userID string, rule // UpdateBudgetRule updates an existing account-level budget rule. func (m *managerImpl) UpdateBudgetRule(ctx context.Context, userID string, rule *types.AccountBudgetRule) (*types.AccountBudgetRule, error) { - if err := m.requirePermission(ctx, rule.AccountID, userID, operations.Update); err != nil { + if err := m.requirePermission(ctx, rule.AccountID, userID, modules.AgentNetworkBudgets, operations.Update); err != nil { return nil, err } @@ -536,7 +541,7 @@ func (m *managerImpl) UpdateBudgetRule(ctx context.Context, userID string, rule // DeleteBudgetRule removes an account-level budget rule. func (m *managerImpl) DeleteBudgetRule(ctx context.Context, accountID, userID, ruleID string) error { - if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkBudgets, operations.Delete); err != nil { return err } @@ -565,7 +570,7 @@ func (m *managerImpl) DeleteBudgetRule(ctx context.Context, accountID, userID, r // emission), a reconcile is triggered so the proxy and peer network maps // converge on the new state. func (m *managerImpl) UpdateSettings(ctx context.Context, userID string, settings *types.Settings) (*types.Settings, error) { - if err := m.requirePermission(ctx, settings.AccountID, userID, operations.Update); err != nil { + if err := m.requirePermission(ctx, settings.AccountID, userID, modules.AgentNetworkSettings, operations.Update); err != nil { return nil, err } @@ -587,6 +592,12 @@ func (m *managerImpl) UpdateSettings(ctx context.Context, userID string, setting if requestedCluster == "" { return status.Errorf(status.NotFound, "agent network settings have not been bootstrapped yet; pass cluster to bootstrap them, or create a provider with bootstrap_cluster set") } + // Bootstrapping pins the cluster and subdomain — a settings + // create on top of the update the caller already passed, matching + // the gate on the provider-create bootstrap path. + if err := m.requirePermission(ctx, settings.AccountID, userID, modules.AgentNetworkSettings, operations.Create); err != nil { + return err + } existing, err = m.bootstrapSettingsIfNeeded(ctx, tx, settings.AccountID, requestedCluster) if err != nil { return err @@ -653,7 +664,7 @@ func (m *managerImpl) validateProviderRefs(ctx context.Context, accountID string // persisting) with cluster and subdomain empty — settings always read as an // object, like the account and DNS settings endpoints. func (m *managerImpl) GetSettings(ctx context.Context, accountID, userID string) (*types.Settings, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkSettings, operations.Read); err != nil { return nil, err } settings, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID) @@ -667,6 +678,21 @@ func (m *managerImpl) GetSettings(ctx context.Context, accountID, userID string) } } +// requireSettingsBootstrapPermission gates the one-time settings bootstrap a +// first provider create performs. Pinning the account's cluster and subdomain +// is a settings write, so it needs the settings permission on top of the +// provider one. No-op once the settings row exists. +func (m *managerImpl) requireSettingsBootstrapPermission(ctx context.Context, accountID, userID string) error { + _, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID) + if err == nil { + return nil + } + if !isNotFound(err) { + return fmt.Errorf("get agent network settings: %w", err) + } + return m.requirePermission(ctx, accountID, userID, modules.AgentNetworkSettings, operations.Create) +} + // bootstrapSettingsIfNeeded creates the per-account agent-network // settings row when missing. The cluster comes from the create-time // hint the dashboard sends (auto-picked from the active cluster list); @@ -725,7 +751,7 @@ func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, st store.St // counter view; permission gate is the same Read role that gates // every other agent-network surface. func (m *managerImpl) ListConsumption(ctx context.Context, accountID, userID string) ([]*types.Consumption, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkUsage, operations.Read); err != nil { return nil, err } return m.store.ListAgentNetworkConsumption(ctx, store.LockingStrengthNone, accountID) @@ -734,7 +760,7 @@ func (m *managerImpl) ListConsumption(ctx context.Context, accountID, userID str // ListAccessLogs returns a paginated, server-side-filtered page of // agent-network access logs plus the total count matching the filter. func (m *managerImpl) ListAccessLogs(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLog, int64, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkLogs, operations.Read); err != nil { return nil, 0, err } return m.store.GetAgentNetworkAccessLogs(ctx, store.LockingStrengthNone, accountID, filter) @@ -744,7 +770,7 @@ func (m *managerImpl) ListAccessLogs(ctx context.Context, accountID, userID stri // agent-network access logs grouped by session, plus the total number of // sessions matching the filter. func (m *managerImpl) ListAccessLogSessions(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLogSession, int64, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkLogs, operations.Read); err != nil { return nil, 0, err } return m.store.GetAgentNetworkAccessLogSessions(ctx, store.LockingStrengthNone, accountID, filter) @@ -753,7 +779,7 @@ func (m *managerImpl) ListAccessLogSessions(ctx context.Context, accountID, user // GetUsageOverview returns the filtered usage rows aggregated into time buckets // at the requested granularity, oldest-first. func (m *managerImpl) GetUsageOverview(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter, granularity types.UsageGranularity) ([]*types.AgentNetworkUsageBucket, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkUsage, operations.Read); err != nil { return nil, err } rows, err := m.store.GetAgentNetworkUsageRows(ctx, store.LockingStrengthNone, accountID, filter) @@ -827,8 +853,8 @@ func (m *managerImpl) RecordConsumption(ctx context.Context, accountID string, k return m.store.IncrementAgentNetworkConsumption(ctx, accountID, kind, dimID, windowSeconds, windowStart, tokensIn, tokensOut, costUSD) } -func (m *managerImpl) requirePermission(ctx context.Context, accountID, userID string, op operations.Operation) error { - ok, _, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.AgentNetwork, op) +func (m *managerImpl) requirePermission(ctx context.Context, accountID, userID string, module modules.Module, op operations.Operation) error { + ok, _, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, module, op) if err != nil { return status.NewPermissionValidationError(err) } diff --git a/management/internals/modules/agentnetwork/provider_bootstrap_test.go b/management/internals/modules/agentnetwork/provider_bootstrap_test.go new file mode 100644 index 000000000..1a2904c51 --- /dev/null +++ b/management/internals/modules/agentnetwork/provider_bootstrap_test.go @@ -0,0 +1,134 @@ +package agentnetwork + +import ( + "context" + "runtime" + "testing" + + "github.com/golang/mock/gomock" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + "github.com/netbirdio/netbird/management/server/account" + "github.com/netbirdio/netbird/management/server/permissions" + "github.com/netbirdio/netbird/management/server/permissions/modules" + "github.com/netbirdio/netbird/management/server/permissions/operations" + "github.com/netbirdio/netbird/management/server/store" + nbtypes "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +// bootstrapFixture wires a real sqlite store to a gomock permissions manager +// so tests can grant the provider permission while denying (or never +// expecting) the settings one. +type bootstrapFixture struct { + manager Manager + store store.Store + perms *permissions.MockManager +} + +func newBootstrapFixture(t *testing.T) *bootstrapFixture { + t.Helper() + if runtime.GOOS == "windows" { + t.Skip("sqlite store not properly supported on Windows yet") + } + t.Setenv("NETBIRD_STORE_ENGINE", string(nbtypes.SqliteStoreEngine)) + + st, cleanUp, err := store.NewTestStoreFromSQL(context.Background(), "", t.TempDir()) + require.NoError(t, err, "test store setup must succeed") + t.Cleanup(cleanUp) + + ctrl := gomock.NewController(t) + perms := permissions.NewMockManager(ctrl) + + accounts := account.NewMockManager(ctrl) + accounts.EXPECT().StoreEvent(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + accounts.EXPECT().UpdateAccountPeers(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + accounts.EXPECT().BufferUpdateAccountPeers(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + + return &bootstrapFixture{ + manager: NewManager(st, perms, accounts, nil), + store: st, + perms: perms, + } +} + +func (f *bootstrapFixture) expectPermission(accountID, userID string, module modules.Module, op operations.Operation, allowed bool) { + f.perms.EXPECT(). + ValidateUserPermissions(gomock.Any(), accountID, userID, module, op). + Return(allowed, context.Background(), nil) +} + +func newBootstrapProvider(accountID string) *types.Provider { + p := types.NewProvider(accountID) + p.Name = "openai" + p.UpstreamURL = "https://api.openai.com" + p.APIKey = "sk-test" + p.Enabled = true + return p +} + +// TestCreateProviderBootstrapRequiresSettingsPermission pins the gate on the +// one-time settings bootstrap: creating the first provider with a +// bootstrap_cluster pins the account's cluster and subdomain, which is a +// settings write and must not ride on the providers permission alone. +func TestCreateProviderBootstrapRequiresSettingsPermission(t *testing.T) { + ctx := context.Background() + + t.Run("denied without settings permission", func(t *testing.T) { + f := newBootstrapFixture(t) + f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true) + f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, false) + + _, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "cluster1.example.com") + require.Error(t, err, "bootstrap without settings permission must fail") + var sErr *status.Error + require.ErrorAs(t, err, &sErr) + assert.Equal(t, status.PermissionDenied, sErr.Type(), "denial should surface as permission denied") + + providers, err := f.store.GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, "account1") + require.NoError(t, err) + assert.Empty(t, providers, "provider must not be persisted when bootstrap is denied") + _, err = f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1") + assert.Error(t, err, "settings row must not be created when bootstrap is denied") + }) + + t.Run("allowed with settings permission", func(t *testing.T) { + f := newBootstrapFixture(t) + f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true) + f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true) + + created, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "cluster1.example.com") + require.NoError(t, err, "bootstrap with both permissions must succeed") + require.NotNil(t, created) + + settings, err := f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1") + require.NoError(t, err, "bootstrap must create the settings row") + assert.Equal(t, "cluster1.example.com", settings.Cluster, "settings should pin the bootstrap cluster") + }) + + t.Run("existing settings need no settings permission", func(t *testing.T) { + f := newBootstrapFixture(t) + require.NoError(t, f.store.SaveAgentNetworkSettings(ctx, &types.Settings{ + AccountID: "account1", + Cluster: "cluster1.example.com", + Subdomain: "existing", + }), "pre-existing settings row setup must succeed") + + // Only the providers permission may be consulted: gomock fails the + // test on any unexpected settings-permission call. + f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true) + + _, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "cluster1.example.com") + require.NoError(t, err, "create with existing settings must not require the settings permission") + }) + + t.Run("no bootstrap cluster needs no settings permission", func(t *testing.T) { + f := newBootstrapFixture(t) + f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true) + + _, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "") + require.NoError(t, err, "create without bootstrap must not require the settings permission") + }) +} diff --git a/management/internals/server/boot.go b/management/internals/server/boot.go index 1c78af9d0..2a1e521aa 100644 --- a/management/internals/server/boot.go +++ b/management/internals/server/boot.go @@ -24,13 +24,13 @@ import ( "github.com/netbirdio/netbird/encryption" "github.com/netbirdio/netbird/formatter/hook" + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" accesslogsmanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs/manager" rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc" "github.com/netbirdio/netbird/management/server/activity" activitystore "github.com/netbirdio/netbird/management/server/activity/store" - "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" nbcache "github.com/netbirdio/netbird/management/server/cache" nbContext "github.com/netbirdio/netbird/management/server/context" nbhttp "github.com/netbirdio/netbird/management/server/http" @@ -184,6 +184,10 @@ func (s *BaseServer) GRPCServer() *grpc.Server { grpc.ChainStreamInterceptor(realip.StreamServerInterceptorOpts(realipOpts...), streamInterceptor, proxyStream), } + // Append interceptors contributed by registered gRPC extensions. These + // run after the built-in chain (ChainUnaryInterceptor is additive). + gRPCOpts = appendExtensionInterceptors(gRPCOpts, s.grpcExtensions) + if s.Config.HttpConfig.LetsEncryptDomain != "" { certManager, err := encryption.CreateCertManager(s.Config.Datadir, s.Config.HttpConfig.LetsEncryptDomain) if err != nil { @@ -215,6 +219,9 @@ func (s *BaseServer) GRPCServer() *grpc.Server { mgmtProto.RegisterProxyServiceServer(gRPCAPIHandler, s.ReverseProxyGRPCServer()) log.Info("ProxyService registered on gRPC server") + // Register services contributed by external modules via the extension seam. + registerExtensions(gRPCAPIHandler, s.grpcExtensions) + return gRPCAPIHandler }) } diff --git a/management/internals/server/grpc_extension.go b/management/internals/server/grpc_extension.go new file mode 100644 index 000000000..3f257c75e --- /dev/null +++ b/management/internals/server/grpc_extension.go @@ -0,0 +1,74 @@ +package server + +import ( + "context" + + "google.golang.org/grpc" +) + +// GRPCExtension bundles an external module's contribution to the management +// gRPC server: the registration of one or more services onto the shared +// grpc.Server, any server-wide interceptors those services require, and an +// optional shutdown hook. It is a generic extension point with no knowledge of +// any specific service. +type GRPCExtension struct { + // Register is invoked with the shared grpc.Server (as a ServiceRegistrar) + // after the built-in services are registered. It may register any number of + // services. May be nil. + Register func(grpc.ServiceRegistrar) + // UnaryInterceptors are appended to the server's unary interceptor chain, + // running after the built-in interceptors. May be empty. + UnaryInterceptors []grpc.UnaryServerInterceptor + // StreamInterceptors are appended to the server's stream interceptor chain, + // running after the built-in interceptors. May be empty. + StreamInterceptors []grpc.StreamServerInterceptor + // Shutdown, if non-nil, is called once during Stop() with the context + // governing server shutdown, which carries a deadline. The hook MUST + // return promptly and MUST abandon its work once that context is + // cancelled or expires: it runs before the rest of Stop()'s cleanup + // (store, event store, embedded IdP) and before Stop() itself checks the + // context's deadline, so a hook that ignores the context will delay all + // of that cleanup and prevent Stop() from returning on time. May be nil. + Shutdown func(ctx context.Context) +} + +// RegisterGRPCExtension registers a gRPC extension. Call before the gRPC server +// is first built (i.e. before Start); registrations after that have no effect. +func (s *BaseServer) RegisterGRPCExtension(ext GRPCExtension) { + s.grpcExtensions = append(s.grpcExtensions, ext) +} + +// appendExtensionInterceptors appends each extension's interceptors to the gRPC +// server options as additional chained interceptors. grpc.ChainUnaryInterceptor +// and grpc.ChainStreamInterceptor are additive, so the returned options run the +// extension interceptors after any interceptors already present in opts. +func appendExtensionInterceptors(opts []grpc.ServerOption, exts []GRPCExtension) []grpc.ServerOption { + for _, ext := range exts { + if len(ext.UnaryInterceptors) > 0 { + opts = append(opts, grpc.ChainUnaryInterceptor(ext.UnaryInterceptors...)) + } + if len(ext.StreamInterceptors) > 0 { + opts = append(opts, grpc.ChainStreamInterceptor(ext.StreamInterceptors...)) + } + } + return opts +} + +// registerExtensions registers each extension's services onto reg. +func registerExtensions(reg grpc.ServiceRegistrar, exts []GRPCExtension) { + for _, ext := range exts { + if ext.Register != nil { + ext.Register(reg) + } + } +} + +// runExtensionShutdownHooks calls each extension's shutdown hook, if set, +// passing ctx through so hooks can honor its deadline/cancellation. +func runExtensionShutdownHooks(ctx context.Context, exts []GRPCExtension) { + for _, ext := range exts { + if ext.Shutdown != nil { + ext.Shutdown(ctx) + } + } +} diff --git a/management/internals/server/grpc_extension_test.go b/management/internals/server/grpc_extension_test.go new file mode 100644 index 000000000..8f444ca72 --- /dev/null +++ b/management/internals/server/grpc_extension_test.go @@ -0,0 +1,160 @@ +package server + +import ( + "context" + "net" + "sync/atomic" + "testing" + + "google.golang.org/grpc" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/grpc/health" + healthgrpc "google.golang.org/grpc/health/grpc_health_v1" + "google.golang.org/grpc/test/bufconn" +) + +// Test that an extension's interceptors and service registration are actually +// wired onto a real in-process gRPC server via the helpers, and that shutdown +// hooks run. This validates the load-bearing assumption that +// grpc.ChainUnaryInterceptor is additive (extension interceptors run in +// addition to any base chain). +func TestGRPCExtensionAppliedToServer(t *testing.T) { + var unaryCalls atomic.Int32 + var streamShutdownCalled atomic.Bool + + ext := GRPCExtension{ + Register: func(reg grpc.ServiceRegistrar) { + healthgrpc.RegisterHealthServer(reg, health.NewServer()) + }, + UnaryInterceptors: []grpc.UnaryServerInterceptor{ + func(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) { + unaryCalls.Add(1) + return handler(ctx, req) + }, + }, + Shutdown: func(ctx context.Context) { streamShutdownCalled.Store(true) }, + } + exts := []GRPCExtension{ext} + + // Base options mimic GRPCServer(): a pre-existing chain the extension appends to. + var baseUnaryCalls atomic.Int32 + opts := []grpc.ServerOption{ + grpc.ChainUnaryInterceptor(func(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) { + baseUnaryCalls.Add(1) + return handler(ctx, req) + }), + } + opts = appendExtensionInterceptors(opts, exts) + + srv := grpc.NewServer(opts...) + registerExtensions(srv, exts) + + lis := bufconn.Listen(1024 * 1024) + go func() { _ = srv.Serve(lis) }() + t.Cleanup(srv.Stop) + + conn, err := grpc.NewClient("passthrough:///bufnet", + grpc.WithContextDialer(func(ctx context.Context, _ string) (net.Conn, error) { return lis.DialContext(ctx) }), + grpc.WithTransportCredentials(insecure.NewCredentials())) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = conn.Close() }) + + _, err = healthgrpc.NewHealthClient(conn).Check(context.Background(), &healthgrpc.HealthCheckRequest{}) + if err != nil { + t.Fatalf("health check via extension-registered service failed: %v", err) + } + if baseUnaryCalls.Load() != 1 { + t.Errorf("base interceptor calls = %d, want 1 (base chain must be preserved)", baseUnaryCalls.Load()) + } + if unaryCalls.Load() != 1 { + t.Errorf("extension interceptor calls = %d, want 1", unaryCalls.Load()) + } + + runExtensionShutdownHooks(context.Background(), exts) + if !streamShutdownCalled.Load() { + t.Error("extension shutdown hook was not called") + } +} + +// TestGRPCExtensionShutdownHookReceivesCallerContext asserts that each hook receives +// a non-nil context and that it is the very same context the caller passed +// in, so hooks can rely on values/deadlines placed on it by Stop(). +func TestGRPCExtensionShutdownHookReceivesCallerContext(t *testing.T) { + type sentinelKey struct{} + want := "shutdown-ctx-sentinel" + ctx := context.WithValue(context.Background(), sentinelKey{}, want) + + var called bool + ext := GRPCExtension{ + Shutdown: func(hookCtx context.Context) { + called = true + if hookCtx == nil { + t.Fatal("hook received a nil context") + } + got, _ := hookCtx.Value(sentinelKey{}).(string) + if got != want { + t.Errorf("hook context sentinel = %q, want %q (not the caller's context)", got, want) + } + }, + } + + runExtensionShutdownHooks(ctx, []GRPCExtension{ext}) + if !called { + t.Fatal("shutdown hook was not called") + } +} + +// TestGRPCExtensionShutdownHookObservesCancellation documents, by test, that +// hooks can honor cancellation/deadlines: a hook given an already-cancelled +// context must see ctx.Err() != nil and a closed Done() channel. +func TestGRPCExtensionShutdownHookObservesCancellation(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + var called bool + ext := GRPCExtension{ + Shutdown: func(hookCtx context.Context) { + called = true + if hookCtx.Err() == nil { + t.Error("hook context Err() = nil, want non-nil for a cancelled context") + } + select { + case <-hookCtx.Done(): + default: + t.Error("hook context Done() channel is not closed for a cancelled context") + } + }, + } + + runExtensionShutdownHooks(ctx, []GRPCExtension{ext}) + if !called { + t.Fatal("shutdown hook was not called") + } +} + +// TestGRPCExtensionShutdownHookNilSkipped asserts that an extension +// with a nil Shutdown hook is skipped without panicking, and that hooks for +// other extensions still run. +func TestGRPCExtensionShutdownHookNilSkipped(t *testing.T) { + var called atomic.Bool + exts := []GRPCExtension{ + {Shutdown: nil}, + {Shutdown: func(context.Context) { called.Store(true) }}, + } + + runExtensionShutdownHooks(context.Background(), exts) + if !called.Load() { + t.Error("shutdown hook for non-nil extension was not called") + } +} + +func TestRegisterGRPCExtensionAccumulates(t *testing.T) { + s := &BaseServer{} + s.RegisterGRPCExtension(GRPCExtension{}) + s.RegisterGRPCExtension(GRPCExtension{}) + if len(s.grpcExtensions) != 2 { + t.Fatalf("grpcExtensions len = %d, want 2", len(s.grpcExtensions)) + } +} diff --git a/management/internals/server/server.go b/management/internals/server/server.go index 7fd06d947..22a61bada 100644 --- a/management/internals/server/server.go +++ b/management/internals/server/server.go @@ -68,6 +68,11 @@ type BaseServer struct { proxyAuthClose func() + // grpcExtensions holds additional gRPC services, interceptors, and shutdown + // hooks registered by external modules via RegisterGRPCExtension. Populated + // during boot (single-threaded), consumed by GRPCServer() and Stop(). + grpcExtensions []GRPCExtension + listener net.Listener certManager *autocert.Manager update *version.Update @@ -257,6 +262,7 @@ func (s *BaseServer) Stop() error { s.proxyAuthClose() s.proxyAuthClose = nil } + runExtensionShutdownHooks(ctx, s.grpcExtensions) _ = s.Store().Close(ctx) _ = s.EventStore().Close(ctx) if s.update != nil { diff --git a/management/internals/shared/grpc/components_encoder.go b/management/internals/shared/grpc/components_encoder.go index 7e43cf478..e1a5cae48 100644 --- a/management/internals/shared/grpc/components_encoder.go +++ b/management/internals/shared/grpc/components_encoder.go @@ -61,6 +61,8 @@ func EncodeNetworkMapEnvelope(in ComponentsEnvelopeInput) *proto.NetworkMapEnvel return &proto.NetworkMapEnvelope{ Payload: &proto.NetworkMapEnvelope_Full{ Full: &proto.NetworkMapComponentsFull{ + Serial: networkSerial(c.Network), + Network: toAccountNetwork(c.Network), PeerConfig: in.PeerConfig, // components.Peers always contains the target peer Peers: []*proto.PeerCompact{toPeerCompact(c.Peers[c.PeerID])}, diff --git a/management/internals/shared/grpc/components_encoder_test.go b/management/internals/shared/grpc/components_encoder_test.go index 100ab0948..f7df82f2f 100644 --- a/management/internals/shared/grpc/components_encoder_test.go +++ b/management/internals/shared/grpc/components_encoder_test.go @@ -758,6 +758,9 @@ func TestEncodeNetworkMapEnvelope_NilComponentsGracefulDegrade(t *testing.T) { assert.Equal(t, "netbird.cloud", full.DnsDomain) assert.Len(t, full.Peers, 1) assert.Empty(t, full.Policies) + require.NotNil(t, full.Network, "client runs Calculate() over the envelope and dereferences Network unconditionally; a nil here would crash the receiver") + assert.Equal(t, "net-empty", full.Network.Identifier) + assert.Equal(t, uint64(9), full.Serial) } func TestEncodeNetworkMapEnvelope_AccountSettingsAlwaysEmitted(t *testing.T) { @@ -776,6 +779,12 @@ func TestEncodeNetworkMapEnvelope_AccountSettingsAlwaysEmitted(t *testing.T) { func emptyNetworkMapComponents() *types.NetworkMapComponents { return types.EmptyNetworkMapComponents( &types.NetworkMapComponents{ - PeerID: "peer-id", Peers: map[string]*types.ComponentPeer{"peer-id": {}}}, + PeerID: "peer-id", Peers: map[string]*types.ComponentPeer{"peer-id": {}}, + Network: &types.Network{ + Identifier: "net-empty", + Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}, + Serial: 9, + }, + }, ) } diff --git a/management/server/permissions/manager.go b/management/server/permissions/manager.go index 995f234d8..6b9977a86 100644 --- a/management/server/permissions/manager.go +++ b/management/server/permissions/manager.go @@ -82,6 +82,9 @@ func (m *managerImpl) ValidateUserPermissions( return m.ValidateRoleModuleAccess(ctx, accountID, role, module, operation), ctxEnriched, nil } +// ValidateRoleModuleAccess resolves an operation against the role's explicit +// grant for the module, then the grant for its parent module when the module +// is a dotted submodule, and finally the role's AutoAllowNew default. func (m *managerImpl) ValidateRoleModuleAccess( ctx context.Context, accountID string, @@ -89,7 +92,7 @@ func (m *managerImpl) ValidateRoleModuleAccess( module modules.Module, operation operations.Operation, ) bool { - if permissions, ok := role.Permissions[module]; ok { + if permissions, ok := lookupModulePermissions(role, module); ok { if allowed, exists := permissions[operation]; exists { return allowed } @@ -100,6 +103,21 @@ func (m *managerImpl) ValidateRoleModuleAccess( return role.AutoAllowNew[operation] } +// lookupModulePermissions returns the role's explicit permission set for the +// module, falling back to the parent module's set for dotted submodules. The +// second return reports whether any explicit set was found. +func lookupModulePermissions(role roles.RolePermissions, module modules.Module) (map[operations.Operation]bool, bool) { + if permissions, ok := role.Permissions[module]; ok { + return permissions, true + } + if parent, hasParent := module.Parent(); hasParent { + if permissions, ok := role.Permissions[parent]; ok { + return permissions, true + } + } + return nil, false +} + func (m *managerImpl) ValidateAccountAccess(ctx context.Context, accountID string, user *types.User, allowOwnerAndAdmin bool) (context.Context, error) { if user.AccountID != accountID { return ctx, status.NewUserNotPartOfAccountError() @@ -119,7 +137,7 @@ func (m *managerImpl) GetPermissionsByRole(ctx context.Context, role types.UserR permissions := roles.Permissions{} for k := range modules.All { - if rolePermissions, ok := roleMap.Permissions[k]; ok { + if rolePermissions, ok := lookupModulePermissions(roleMap, k); ok { permissions[k] = rolePermissions continue } diff --git a/management/server/permissions/manager_test.go b/management/server/permissions/manager_test.go new file mode 100644 index 000000000..345212f43 --- /dev/null +++ b/management/server/permissions/manager_test.go @@ -0,0 +1,139 @@ +package permissions + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/server/permissions/modules" + "github.com/netbirdio/netbird/management/server/permissions/operations" + "github.com/netbirdio/netbird/management/server/permissions/roles" + "github.com/netbirdio/netbird/management/server/types" +) + +func TestValidateRoleModuleAccessSubmoduleCascade(t *testing.T) { + manager := NewManager(nil) + ctx := context.Background() + + fullAccess := map[operations.Operation]bool{ + operations.Read: true, + operations.Create: true, + operations.Update: true, + operations.Delete: true, + } + readOnly := map[operations.Operation]bool{ + operations.Read: true, + operations.Create: false, + operations.Update: false, + operations.Delete: false, + } + denyAll := map[operations.Operation]bool{ + operations.Read: false, + operations.Create: false, + operations.Update: false, + operations.Delete: false, + } + + t.Run("parent grant covers submodules", func(t *testing.T) { + role := roles.RolePermissions{ + AutoAllowNew: denyAll, + Permissions: roles.Permissions{modules.AgentNetwork: fullAccess}, + } + assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Create), + "parent full grant should allow create on a submodule") + assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkLogs, operations.Read), + "parent full grant should allow read on a submodule") + }) + + t.Run("submodule grant does not leak to parent or siblings", func(t *testing.T) { + role := roles.RolePermissions{ + AutoAllowNew: denyAll, + Permissions: roles.Permissions{modules.AgentNetworkUsage: readOnly}, + } + assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkUsage, operations.Read), + "explicit submodule read should be allowed") + assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkUsage, operations.Create), + "read-only submodule grant should not allow create") + assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetwork, operations.Read), + "submodule grant should not grant the parent module") + assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Read), + "submodule grant should not grant a sibling submodule") + }) + + t.Run("explicit submodule entry wins over parent grant", func(t *testing.T) { + role := roles.RolePermissions{ + AutoAllowNew: denyAll, + Permissions: roles.Permissions{ + modules.AgentNetwork: fullAccess, + modules.AgentNetworkLogs: denyAll, + }, + } + assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkLogs, operations.Read), + "explicit submodule deny should override the parent grant") + assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkUsage, operations.Read), + "sibling submodules should still resolve through the parent grant") + }) + + t.Run("auto allow applies when neither submodule nor parent is granted", func(t *testing.T) { + role := roles.RolePermissions{ + AutoAllowNew: readOnly, + } + assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Read), + "auto-allow read should apply to submodules") + assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Delete), + "auto-allow should not grant unlisted operations") + }) +} + +// TestExistingRolesKeepAgentNetworkBehaviorOnSubmodules pins the behavior the +// submodule split must not change: every built-in role resolves the new +// submodules exactly as it resolved the agent_network module before. +func TestExistingRolesKeepAgentNetworkBehaviorOnSubmodules(t *testing.T) { + manager := NewManager(nil) + ctx := context.Background() + + submodules := []modules.Module{ + modules.AgentNetworkProviders, + modules.AgentNetworkPolicies, + modules.AgentNetworkGuardrails, + modules.AgentNetworkBudgets, + modules.AgentNetworkUsage, + modules.AgentNetworkLogs, + modules.AgentNetworkSettings, + } + allOperations := []operations.Operation{operations.Read, operations.Create, operations.Update, operations.Delete} + + for _, role := range []types.UserRole{types.UserRoleOwner, types.UserRoleAdmin, types.UserRoleAuditor, types.UserRoleNetworkAdmin, types.UserRoleUser} { + rolePermissions, ok := roles.RolesMap[role] + require.True(t, ok, "role %s must exist in RolesMap", role) + + for _, sub := range submodules { + for _, op := range allOperations { + expected := manager.ValidateRoleModuleAccess(ctx, "account", rolePermissions, modules.AgentNetwork, op) + actual := manager.ValidateRoleModuleAccess(ctx, "account", rolePermissions, sub, op) + assert.Equal(t, expected, actual, "role %s: %s on %s should match the agent_network module", role, op, sub) + } + } + } +} + +func TestGetPermissionsByRoleIncludesSubmodules(t *testing.T) { + manager := NewManager(nil) + ctx := context.Background() + + permissions, err := manager.GetPermissionsByRole(ctx, types.UserRoleAuditor) + require.NoError(t, err, "auditor role must resolve") + + usage, ok := permissions[modules.AgentNetworkUsage] + require.True(t, ok, "permissions map should contain the usage submodule") + assert.True(t, usage[operations.Read], "auditor should read the usage submodule") + assert.False(t, usage[operations.Update], "auditor should not update the usage submodule") + + adminPermissions, err := manager.GetPermissionsByRole(ctx, types.UserRoleAdmin) + require.NoError(t, err, "admin role must resolve") + providers, ok := adminPermissions[modules.AgentNetworkProviders] + require.True(t, ok, "permissions map should contain the providers submodule") + assert.True(t, providers[operations.Delete], "admin should delete on the providers submodule") +} diff --git a/management/server/permissions/modules/module.go b/management/server/permissions/modules/module.go index a3a9c554d..8a2a1a52d 100644 --- a/management/server/permissions/modules/module.go +++ b/management/server/permissions/modules/module.go @@ -1,5 +1,7 @@ package modules +import "strings" + type Module string const ( @@ -20,6 +22,17 @@ const ( IdentityProviders Module = "identity_providers" Services Module = "services" AgentNetwork Module = "agent_network" + + // Agent Network submodules. A role may grant one of these directly + // or grant the AgentNetwork parent, which covers all of them (see + // permissions.Manager cascade resolution). + AgentNetworkProviders Module = "agent_network.providers" + AgentNetworkPolicies Module = "agent_network.policies" + AgentNetworkGuardrails Module = "agent_network.guardrails" + AgentNetworkBudgets Module = "agent_network.budgets" + AgentNetworkUsage Module = "agent_network.usage" + AgentNetworkLogs Module = "agent_network.logs" + AgentNetworkSettings Module = "agent_network.settings" ) var All = map[Module]struct{}{ @@ -40,4 +53,21 @@ var All = map[Module]struct{}{ IdentityProviders: {}, Services: {}, AgentNetwork: {}, + + AgentNetworkProviders: {}, + AgentNetworkPolicies: {}, + AgentNetworkGuardrails: {}, + AgentNetworkBudgets: {}, + AgentNetworkUsage: {}, + AgentNetworkLogs: {}, + AgentNetworkSettings: {}, +} + +// Parent returns the module owning a dotted submodule name and true, or the +// module itself and false when it has no parent. +func (m Module) Parent() (Module, bool) { + if i := strings.IndexByte(string(m), '.'); i > 0 { + return Module(string(m)[:i]), true + } + return m, false } diff --git a/management/server/types/proxy_access_token.go b/management/server/types/proxy_access_token.go index b20b83bc1..9bb27ef02 100644 --- a/management/server/types/proxy_access_token.go +++ b/management/server/types/proxy_access_token.go @@ -68,7 +68,7 @@ type ProxyAccessTokenGenerated struct { // CreateNewProxyAccessToken generates a new proxy access token. // Returns the token with hashed value stored and plain token for one-time display. func CreateNewProxyAccessToken(name string, expiresIn time.Duration, accountID *string, createdBy string) (*ProxyAccessTokenGenerated, error) { - hashedToken, plainToken, err := generateProxyToken() + hashedToken, plainToken, err := GenerateProxyToken() if err != nil { return nil, err } @@ -94,7 +94,10 @@ func CreateNewProxyAccessToken(name string, expiresIn time.Duration, accountID * }, nil } -func generateProxyToken() (HashedProxyToken, PlainProxyToken, error) { +// GenerateProxyToken generates a new random proxy token, returning its SHA-256 +// hash (for storage) and the one-time plaintext. Exported so external modules +// can mint tokens in the canonical proxy-token format. +func GenerateProxyToken() (HashedProxyToken, PlainProxyToken, error) { secret, err := b.Random(ProxyTokenSecretLength) if err != nil { return "", "", err diff --git a/management/server/types/proxy_access_token_test.go b/management/server/types/proxy_access_token_test.go index aa1a4d2dd..740b87c2f 100644 --- a/management/server/types/proxy_access_token_test.go +++ b/management/server/types/proxy_access_token_test.go @@ -1,6 +1,7 @@ package types import ( + "strings" "testing" "time" @@ -123,6 +124,22 @@ func TestCreateNewProxyAccessToken(t *testing.T) { }) } +func TestGenerateProxyToken(t *testing.T) { + hashed, plain, err := GenerateProxyToken() + if err != nil { + t.Fatal(err) + } + if err := plain.Validate(); err != nil { + t.Errorf("generated token failed Validate(): %v", err) + } + if plain.Hash() != hashed { + t.Error("returned hashed token does not match Hash(plain)") + } + if !strings.HasPrefix(string(plain), ProxyTokenPrefix) { + t.Errorf("token %q missing prefix %q", plain, ProxyTokenPrefix) + } +} + func TestProxyAccessToken_IsExpired(t *testing.T) { past := time.Now().Add(-1 * time.Hour) future := time.Now().Add(1 * time.Hour) diff --git a/release_files/darwin_pkg/postinstall b/release_files/darwin_pkg/postinstall index 33fa4bfee..2c96a80cd 100755 --- a/release_files/darwin_pkg/postinstall +++ b/release_files/darwin_pkg/postinstall @@ -30,7 +30,23 @@ mkdir -p /usr/local/bin/ $AGENT service install || true $AGENT service start || true - open $APP + console_user=$(stat -f%Su /dev/console 2>/dev/null) + case "$console_user" in + ""|root|loginwindow|_mbsetupuser) + echo "No active GUI user session (console user: '${console_user:-none}'); skipping UI launch." + ;; + *) + uid=$(id -u "$console_user" 2>/dev/null) + if [ -z "$uid" ]; then + echo "Could not resolve uid for console user '$console_user'; skipping UI launch." + else + echo "Launching NetBird UI as console user $console_user (uid $uid)." + if ! launchctl asuser "$uid" sudo -u "$console_user" -H open "$APP"; then + echo "Failed to launch NetBird UI; if autostart is enabled it will start at next login." + fi + fi + ;; + esac echo "Finished Netbird installation successfully" exit 0 # all good diff --git a/shared/management/networkmap/decode.go b/shared/management/networkmap/decode.go index d15117b6e..4864a9dff 100644 --- a/shared/management/networkmap/decode.go +++ b/shared/management/networkmap/decode.go @@ -228,15 +228,17 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents, return c, nil } +// decodeAccountNetwork never returns nil — Calculate() dereferences +// c.Network unconditionally, and servers that predate the fix omit the field +// entirely from the empty-components envelope. func decodeAccountNetwork(an *proto.AccountNetwork) *types.Network { + n := &types.Network{} if an == nil { - return nil - } - n := &types.Network{ - Identifier: an.Identifier, - Dns: an.Dns, - Serial: an.Serial, + return n } + n.Identifier = an.Identifier + n.Dns = an.Dns + n.Serial = an.Serial if an.NetCidr != "" { if _, ipnet, err := net.ParseCIDR(an.NetCidr); err == nil && ipnet != nil { n.Net = *ipnet diff --git a/shared/management/networkmap/envelope_test.go b/shared/management/networkmap/envelope_test.go index a81478aff..92e1916da 100644 --- a/shared/management/networkmap/envelope_test.go +++ b/shared/management/networkmap/envelope_test.go @@ -221,6 +221,66 @@ func TestEnvelopeRoundTrip_AllGroupShortCircuitParity(t *testing.T) { "client-side Calculate must connect the same remote peers as the server") } +// TestEnvelopeToNetworkMap_EmptyComponents covers the graceful-degrade path +// the server takes for a peer that is missing from the account or absent from +// the validated-peers map. The legacy server short-circuited before +// Calculate() and shipped a NetworkMap carrying only the account Network; the +// components path runs Calculate() on the client instead, so the envelope must +// carry Network or the client panics dereferencing a nil *types.Network. +func TestEnvelopeToNetworkMap_EmptyComponents(t *testing.T) { + localPeerKey := randomWgKey(t) + c := types.EmptyNetworkMapComponents(&types.NetworkMapComponents{ + PeerID: "peer-A", + Network: &types.Network{ + Identifier: "net-empty", + Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}, + Serial: 7, + }, + Peers: map[string]*types.ComponentPeer{ + "peer-A": {ID: "peer-A", Key: localPeerKey, IP: netip.AddrFrom4([4]byte{100, 64, 0, 1})}, + }, + }) + + envelope := mgmtgrpc.EncodeNetworkMapEnvelope(mgmtgrpc.ComponentsEnvelopeInput{ + Components: c, + DNSDomain: "netbird.cloud", + }) + require.NotNil(t, envelope.GetFull().Network, "empty envelope must carry the account Network") + + wire, err := goproto.Marshal(envelope) + require.NoError(t, err, "marshal envelope") + var decoded proto.NetworkMapEnvelope + require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope") + + result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud") + require.NoError(t, err, "EnvelopeToNetworkMap must degrade gracefully on empty components") + require.Equal(t, uint64(7), result.NetworkMap.Serial) + require.Empty(t, result.NetworkMap.RemotePeers, "unvalidated peer connects to nobody") +} + +// TestEnvelopeToNetworkMap_MissingNetwork simulates a server that omits +// AccountNetwork from the envelope. Clients must degrade rather than panic, so +// they survive talking to a management server that predates the encoder fix. +func TestEnvelopeToNetworkMap_MissingNetwork(t *testing.T) { + c, localPeerKey := buildSmokeComponents(t) + + envelope := mgmtgrpc.EncodeNetworkMapEnvelope(mgmtgrpc.ComponentsEnvelopeInput{ + Components: c, + DNSDomain: "netbird.cloud", + }) + envelope.GetFull().Network = nil + + wire, err := goproto.Marshal(envelope) + require.NoError(t, err, "marshal envelope") + var decoded proto.NetworkMapEnvelope + require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope") + + result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud") + require.NoError(t, err, "a missing AccountNetwork must not panic the client") + require.NotNil(t, result.Components.Network) + require.NotEmpty(t, result.NetworkMap.RemotePeers, "the rest of the snapshot stays usable") +} + // buildSmokeComponents returns a minimal NetworkMapComponents (2 peers, 1 // group, 1 allow policy) plus the receiving peer's WG public key. Sufficient // to validate the encode → marshal → decode → Calculate pipeline produces