diff --git a/client/server/login_account_test.go b/client/server/login_account_test.go index c113c8bfe..6198a00bc 100644 --- a/client/server/login_account_test.go +++ b/client/server/login_account_test.go @@ -2,6 +2,8 @@ package server import ( "context" + "errors" + "path/filepath" "testing" "time" @@ -9,9 +11,24 @@ import ( "github.com/netbirdio/netbird/client/internal" "github.com/netbirdio/netbird/client/internal/auth" + "github.com/netbirdio/netbird/client/internal/ipcauth" + "github.com/netbirdio/netbird/client/internal/profilemanager" + "github.com/netbirdio/netbird/client/mdm" "github.com/netbirdio/netbird/client/proto" ) +type jwtStoringMDMFetcher struct { + cache *jwtCache + caller ipcauth.Identity + generation uint64 +} + +func (f *jwtStoringMDMFetcher) Fetch() map[string]any { + f.generation = f.cache.currentGeneration() + f.cache.store("previous-profile-token", f.caller, time.Minute, f.generation) + return nil +} + type stubOAuthFlow struct { token auth.TokenInfo onWait func() @@ -123,6 +140,110 @@ func TestSwitchProfile_DropsAccountPromptAndPendingFlow(t *testing.T) { require.False(t, pending, "the previous profile's extend flow leaked across a profile switch") } +func TestLogin_ProfileSwitchDropsAccountPromptAndPendingFlow(t *testing.T) { + s, _, _, username, cfgPath := setupServerWithProfile(t) + s.rootCtx = internal.CtxInitState(context.Background()) + s.isLoginRequiredFn = func(context.Context) (bool, error) { + return true, nil + } + + other := "other-profile" + otherPath := filepath.Join(filepath.Dir(cfgPath), other+".json") + _, err := profilemanager.UpdateOrCreateConfig(profilemanager.ConfigInput{ + ConfigPath: otherPath, + ManagementURL: "https://api.netbird.io:443", + }) + require.NoError(t, err) + breakProfilePrivateKey(t, otherPath) + + s.forceAccountPrompt = true + cancelled := false + s.oauthAuthFlow = oauthAuthFlow{ + flow: &stubOAuthFlow{}, + hint: "user@example.com", + waitCancel: func() { cancelled = true }, + } + + extendCancelled := false + s.extendAuthSessionFlow.Set(&stubOAuthFlow{}, auth.AuthFlowInfo{DeviceCode: "device"}) + s.extendAuthSessionFlow.SetWaitCancel(func() { extendCancelled = true }) + + generation := s.jwtCache.currentGeneration() + + _, err = s.Login(userCtx(), &proto.LoginRequest{ProfileName: &other, Username: &username}) + require.Error(t, err, "the broken key must stop the login before a flow is built") + + active, err := s.profileManager.GetActiveProfileState() + require.NoError(t, err) + require.Equal(t, profilemanager.ID(other), active.ID, "the login did not switch the profile") + + require.False(t, s.forceAccountPrompt, "the prompt flag leaked across a login-driven profile switch") + require.Nil(t, s.oauthAuthFlow.flow, "the previous profile's flow leaked across a login-driven profile switch") + require.Empty(t, s.oauthAuthFlow.hint) + require.True(t, cancelled, "the pending wait was not cancelled") + + require.True(t, extendCancelled, "the pending extend wait was not cancelled") + _, _, pending := s.extendAuthSessionFlow.Get() + require.False(t, pending, "the previous profile's extend flow leaked across a login-driven profile switch") + + require.Greater(t, s.jwtCache.currentGeneration(), generation, "the previous profile's JWT cache survived a login-driven profile switch") +} + +func TestLogin_ProfileSwitchRejectsJWTObtainedUnderPreviousConfig(t *testing.T) { + s, _, _, username, cfgPath := setupServerWithProfile(t) + s.rootCtx = internal.CtxInitState(context.Background()) + s.isLoginRequiredFn = func(context.Context) (bool, error) { + return false, errors.New("stop once the config is swapped") + } + + other := "other-profile" + _, err := profilemanager.UpdateOrCreateConfig(profilemanager.ConfigInput{ + ConfigPath: filepath.Join(filepath.Dir(cfgPath), other+".json"), + ManagementURL: "https://api.netbird.io:443", + }) + require.NoError(t, err) + + // getConfig loads the MDM policy right before Login swaps s.config, so the + // fetcher runs where a RequestJWTAuth racing the login would land: after the + // switch dropped the previous profile's state, while s.config still belongs + // to the previous profile. It caches a token under the generation current at + // that point, as WaitJWTToken would. + generation := s.jwtCache.currentGeneration() + fetcher := &jwtStoringMDMFetcher{cache: s.jwtCache, caller: unprivilegedIdentity()} + s.mdmLoader = mdm.NewLoader(fetcher) + + _, err = s.Login(userCtx(), &proto.LoginRequest{ProfileName: &other, Username: &username}) + require.Error(t, err) + + require.Greater(t, fetcher.generation, generation, "the token was not cached after the switch dropped the previous profile's state") + _, found := s.jwtCache.get(fetcher.caller) + require.False(t, found, "a JWT obtained under the previous profile's config survived the login-driven switch") +} + +func TestLogin_SameProfileKeepsPendingFlow(t *testing.T) { + s, _, profName, username, cfgPath := setupServerWithProfile(t) + s.rootCtx = internal.CtxInitState(context.Background()) + s.isLoginRequiredFn = func(context.Context) (bool, error) { + return true, nil + } + breakProfilePrivateKey(t, cfgPath) + + cancelled := false + flow := &stubOAuthFlow{} + s.oauthAuthFlow = oauthAuthFlow{ + flow: flow, + hint: "user@example.com", + waitCancel: func() { cancelled = true }, + } + + _, err := s.Login(userCtx(), &proto.LoginRequest{ProfileName: &profName, Username: &username}) + require.Error(t, err, "the broken key must stop the login before a flow is built") + + require.Equal(t, flow, s.oauthAuthFlow.flow, "a login on the same profile dropped the flow a second client could join") + require.Equal(t, "user@example.com", s.oauthAuthFlow.hint) + require.False(t, cancelled, "a login on the same profile cancelled the pending wait") +} + func TestWaitSSOLogin_JudgesTheFlowThatProducedTheToken(t *testing.T) { s := newSSOTestServer(t, "user@example.com", false, "user@example.com") attempts := 0 diff --git a/client/server/server.go b/client/server/server.go index 916d7041b..e082e1f2a 100644 --- a/client/server/server.go +++ b/client/server/server.go @@ -722,7 +722,7 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro } }() - ctx, activeProf, err := s.authorizeAndPrepareLogin(callerCtx, msg, activeProf) + ctx, activeProf, switched, err := s.authorizeAndPrepareLogin(callerCtx, msg, activeProf) if err != nil { // The RPC boundary is where this gets recorded: nothing logs handler // errors for us, and a caller that retries would otherwise leave no @@ -752,6 +752,9 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro } s.mutex.Lock() s.config = config + if switched { + s.jwtCache.clear() + } s.mutex.Unlock() s.localMetrics.Reconcile(config.LocalMetricsEnabled, config.LocalMetricsAddress) @@ -1192,11 +1195,15 @@ func (s *Server) Up(callerCtx context.Context, msg *proto.UpRequest) (*proto.UpR } if msg != nil && msg.ProfileName != nil { - if _, err := s.switchProfileIfNeeded(*msg.ProfileName, msg.Username, activeProf); err != nil { + switched, err := s.switchProfileIfNeeded(*msg.ProfileName, msg.Username, activeProf) + if err != nil { s.mutex.Unlock() log.Errorf("failed to switch profile: %v", err) return nil, err } + if switched { + s.dropPendingAuthFlows() + } } activeProf, err = s.profileManager.GetActiveProfileState() @@ -1334,12 +1341,12 @@ func (s *Server) resolveProfileHandle(handle, username string) (*profilemanager. } // switchProfileIfNeeded resolves the user-supplied handle, updates the -// active profile state if it differs from the current one, and returns -// the resolved profile so callers can include its ID in RPC responses. -func (s *Server) switchProfileIfNeeded(handle string, userName *string, activeProf *profilemanager.ActiveProfileState) (*profilemanager.Profile, error) { +// active profile state if it differs from the current one, and reports +// whether the active profile changed. +func (s *Server) switchProfileIfNeeded(handle string, userName *string, activeProf *profilemanager.ActiveProfileState) (bool, error) { if handle != profilemanager.DefaultProfileName && (userName == nil || *userName == "") { log.Errorf("profile name is set to %s, but username is not provided", handle) - return nil, fmt.Errorf("profile name is set to %s, but username is not provided", handle) + return false, fmt.Errorf("profile name is set to %s, but username is not provided", handle) } var username string @@ -1349,26 +1356,48 @@ func (s *Server) switchProfileIfNeeded(handle string, userName *string, activePr resolved, err := s.resolveProfileHandle(handle, username) if err != nil { - return nil, err + return false, err } - if resolved.ID != activeProf.ID || username != activeProf.Username { - if s.checkProfilesDisabled() { - log.Errorf("profiles are disabled, you cannot use this feature without profiles enabled") - return nil, gstatus.Errorf(codes.Unavailable, errProfilesDisabled) - } - - log.Infof("switching to profile %s (%s) for user %s", resolved.Name, resolved.ID, username) - if err := s.profileManager.SetActiveProfileState(&profilemanager.ActiveProfileState{ - ID: resolved.ID, - Username: username, - }); err != nil { - log.Errorf("failed to set active profile state: %v", err) - return nil, fmt.Errorf("failed to set active profile state: %w", err) - } + if resolved.ID == activeProf.ID && username == activeProf.Username { + return false, nil } - return resolved, nil + if s.checkProfilesDisabled() { + log.Errorf("profiles are disabled, you cannot use this feature without profiles enabled") + return false, gstatus.Errorf(codes.Unavailable, errProfilesDisabled) + } + + log.Infof("switching to profile %s (%s) for user %s", resolved.Name, resolved.ID, username) + if err := s.profileManager.SetActiveProfileState(&profilemanager.ActiveProfileState{ + ID: resolved.ID, + Username: username, + }); err != nil { + log.Errorf("failed to set active profile state: %v", err) + return false, fmt.Errorf("failed to set active profile state: %w", err) + } + + return true, nil +} + +func (s *Server) dropPendingAuthFlows() { + // A pending login flow and the account-prompt flag describe the previous + // profile's login; carried across a switch they would judge the new + // profile's token against the old profile's account. CancelFunc is + // non-blocking, so calling it under the mutex is safe. + if cancel := s.oauthAuthFlow.waitCancel; cancel != nil { + cancel() + } + s.oauthAuthFlow = oauthAuthFlow{} + s.forceAccountPrompt = false + + // A pending session extend belongs to the previous profile too: its device + // code was issued by that profile's IdP client, and WaitExtendAuthSession + // would submit the resulting token against the new profile's engine. + s.extendAuthSessionFlow.CancelWait() + s.extendAuthSessionFlow.Clear() + + s.jwtCache.clear() } // SwitchProfile switches the active profile in the daemon. @@ -1402,23 +1431,7 @@ func (s *Server) SwitchProfile(callerCtx context.Context, msg *proto.SwitchProfi s.config = config s.localMetrics.Reconcile(config.LocalMetricsEnabled, config.LocalMetricsAddress) - s.jwtCache.clear() - - // A pending login flow and the account-prompt flag describe the previous - // profile's login; carried across a switch they would judge the new - // profile's token against the old profile's account. CancelFunc is - // non-blocking, so calling it under the mutex is safe. - if cancel := s.oauthAuthFlow.waitCancel; cancel != nil { - cancel() - } - s.oauthAuthFlow = oauthAuthFlow{} - s.forceAccountPrompt = false - - // A pending session extend belongs to the previous profile too: its device - // code was issued by that profile's IdP client, and WaitExtendAuthSession - // would submit the resulting token against the new profile's engine. - s.extendAuthSessionFlow.CancelWait() - s.extendAuthSessionFlow.Clear() + s.dropPendingAuthFlows() if msg != nil && msg.ProfileName != nil { s.publishProfileListChanged(*msg.ProfileName) @@ -2862,7 +2875,7 @@ var afterLoginPreCheck func() // of this is reached; this one exists because that check is not synchronized // against a concurrent privileged request that enables the SSH server, and a // caller refused here must not have cancelled or switched anything either. -func (s *Server) authorizeAndPrepareLogin(callerCtx context.Context, msg *proto.LoginRequest, activeProf *profilemanager.ActiveProfileState) (context.Context, *profilemanager.ActiveProfileState, error) { +func (s *Server) authorizeAndPrepareLogin(callerCtx context.Context, msg *proto.LoginRequest, activeProf *profilemanager.ActiveProfileState) (context.Context, *profilemanager.ActiveProfileState, bool, error) { if afterLoginPreCheck != nil { afterLoginPreCheck() } @@ -2872,10 +2885,10 @@ func (s *Server) authorizeAndPrepareLogin(callerCtx context.Context, msg *proto. stored, err := s.storedLoginConfig(activeProf, msg) if err != nil { - return nil, nil, err + return nil, nil, false, err } if err := requirePrivilegeForConfigChange(callerCtx, stored, privilegedChangeFromLogin(msg)); err != nil { - return nil, nil, err + return nil, nil, false, err } // The update-settings decision is re-taken here for the same reason as the @@ -2884,7 +2897,7 @@ func (s *Server) authorizeAndPrepareLogin(callerCtx context.Context, msg *proto. // authoritative check, and it is the last read before persistLoginOverrides // writes. if s.checkUpdateSettingsDisabled() && configChangeRequested(stored, loginOverridesInput(msg)) { - return nil, nil, gstatus.Errorf(codes.FailedPrecondition, errUpdateSettingsDisabled) + return nil, nil, false, gstatus.Errorf(codes.FailedPrecondition, errUpdateSettingsDisabled) } s.mutex.Lock() @@ -2902,19 +2915,26 @@ func (s *Server) authorizeAndPrepareLogin(callerCtx context.Context, msg *proto. log.Warnf(errRestoreResidualState, err) } + switched := false if msg.ProfileName != nil { - if _, err := s.switchProfileIfNeeded(*msg.ProfileName, msg.Username, activeProf); err != nil { - return nil, nil, fmt.Errorf("switch profile: %w", err) + switched, err = s.switchProfileIfNeeded(*msg.ProfileName, msg.Username, activeProf) + if err != nil { + return nil, nil, false, fmt.Errorf("switch profile: %w", err) + } + if switched { + s.mutex.Lock() + s.dropPendingAuthFlows() + s.mutex.Unlock() } } activeProf, err = s.profileManager.GetActiveProfileState() if err != nil { - return nil, nil, fmt.Errorf("active profile state: %w", err) + return nil, nil, false, fmt.Errorf("active profile state: %w", err) } if err := persistLoginOverrides(activeProf, msg); err != nil { - return nil, nil, fmt.Errorf("persist login overrides: %w", err) + return nil, nil, false, fmt.Errorf("persist login overrides: %w", err) } // Provisioning under the same lock as the decision above, and next to the @@ -2923,10 +2943,10 @@ func (s *Server) authorizeAndPrepareLogin(callerCtx context.Context, msg *proto. // had already answered its caller would be overwritten by the config this // login read before it landed. if _, _, err := provisionProfileIdentity(activeProf); err != nil { - return nil, nil, err + return nil, nil, false, err } - return ctx, activeProf, nil + return ctx, activeProf, switched, nil } // persistLoginOverrides writes the config fields a login request is allowed to