From fe0e9042add59e4ede9c70c18f909d20c5087869 Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Fri, 9 Oct 2026 12:56:41 +0200 Subject: [PATCH 1/2] [client] Drop the pending login flow when a login switches the profile (#7882) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * [client] Force interactive login when extending the auth session A session extend must be answered from the account the peer is registered under. With a silent PKCE flow (DisablePromptLogin or max_age=0) the IdP answers from whatever session it already holds, which need not be the peer's account when several are signed in; the token then fails the user match in ExtendAuthSession with no way to pick another account. Mark the PKCE flow request as a session extend so the management server can force prompt=login for it, overriding the configured silent flow. * [client] Reduce cognitive complexity of Server.Login Login sat at cognitive complexity 27, over the 25 the linter allows. Extract the interactive SSO branch into startSSOLogin, and split the nested in-flight-flow reuse check out of it into reuseOAuthFlow, which flattens the original if/else into early returns: it returns the cached auth info when the previous flow targets the same client and still has more than 90s left, otherwise cancels the stale wait and returns nil so the caller requests a fresh flow. The helpers take the contextState through a small statusSetter interface, since internal.contextState is unexported and re-deriving it with CtxGetState inside the helper would resolve against callerCtx rather than rootCtx. No behavior change: same ordering of state transitions, same mutex scope around the oauthAuthFlow write, same error paths. Login is now at 21. * [client] Respect DisablePromptLogin when extending the auth session Forcing prompt=login on a session extend overrode DisablePromptLogin, which is set for IdPs that break on it: Authentik triggers a double authentication and social logins fail outright. Overriding it there trades a recoverable extend for a login that cannot complete at all. Keep the LoginFlag override, which only replaces max_age=0 or none with prompt=login so the IdP honours login_hint, and leave DisablePromptLogin as configured. Those deployments keep the silent flow, and with several accounts signed in an extend answered from the wrong one still fails the user match. * [client] Guard the shared OAuth flow state with the server mutex reuseOAuthFlow read flow, expiresAt, waitCancel and info without holding s.mutex, while startSSOLogin and WaitSSOLogin write them under it. Reading the fields one at a time could also answer with auth info from a flow that was already replaced, or cancel a wait that no longer belongs to the flow just judged stale. Take one snapshot under the lock and decide from it. WaitSSOLogin read oauthAuthFlow.flow twice outside the lock; both now use a value snapshotted in the critical section that already installs actCancel. Its stale waitCancel was read and called in a separate section from the one installing the new one, so two racing calls could read the same predecessor and leave one wait uncancelled. Swap the two in a single critical section. Both cancels run after unlocking: the displaced wait takes s.mutex as it unwinds. * [client] Verify the SSO login came back for the hinted account login_hint is a suggestion the IdP may ignore: with a silent flow configured (DisablePromptLogin or max_age=0) and a live IdP session for another account, the login completes with that account's token. On a registered peer the management server rejects it as a user mismatch, but on a fresh profile the peer silently registers under the wrong account and the profile is then bound to it — every later login follows the stored hint straight back. After the token exchange, compare the ID token's email against the hint the flow was sent with. On a mismatch, do not log in to management with the token; run one more round asking the IdP to re-decide the account (prompt=login, via ForceAccountPrompt — DisablePromptLogin still wins there). If the prompted round also comes back different, proceed with a warning: the address may legitimately have changed, and refusing forever would lock the user out of the profile while the management server still rejects a token that does not own the peer. A token or profile with no email to compare is not judged. The retry differs per platform because of who opens the browser: - CLI (netbird login foreground) and Android run the whole flow in one process, so the mismatch retries automatically: the browser reopens with the account prompt within the same login attempt. - On desktop the login is split between the daemon and the GUI: Login hands the authorize URL to the GUI, WaitSSOLogin blocks for the token, and only the GUI can open a browser. A new URL cannot be handed out from inside WaitSSOLogin (its response has no field for one, kept that way to avoid a proto change), so the daemon arms forceAccountPrompt, fails the round with "connect again to choose the account", and builds the next Login's flow with the prompt — the user's next connect is the retry. The flag and the flow annotations live in daemon memory only; SwitchProfile drops them so the previous profile's hint cannot judge the next profile's token. The device code flow has no prompt parameter (RFC 8628), so a prompted round there runs as-is and a repeated mismatch is let through with the warning rather than looping. * [client] Address review comments on PKCE session extend flow Fail the PKCE authorization flow test on request error instead of continuing into a nil dereference, and make the godoc comments on the touched exported symbols identifier-leading full sentences. * [client] Match accounts only on the email claim of the ID token The name-claim fallback in the ID token parsing is kept for the login hint and display, but account matching now only considers a value that came from the email claim, so a token without one no longer produces a false account mismatch. * [client] Drop the pending session extend on a profile switch The profile-switch cleanup dropped the pending login flow and the account-prompt flag, but left extendAuthSessionFlow untouched. Its device code was issued by the previous profile's IdP client, so a WaitExtendAuthSession still parked on the browser leg would submit the resulting token against the new profile's engine. * [client] Judge the SSO account against the flow that produced the token WaitSSOLogin snapshotted the flow on entry but re-read the info, hint and accountPrompted from the live s.oauthAuthFlow afterwards, in separate critical sections. WaitToken blocks for the whole browser leg, so a concurrent Login or RequestJWTAuth could replace the flow meanwhile and the mismatch check would compare this wait's token against another flow's account: either arming the prompt spuriously or letting a wrong-account token through against an unrelated profile's hint. Take all of it in the entry snapshot. * [client] Keep the forced account prompt from being lost to flow reuse startSSOLogin consumed forceAccountPrompt and applied the prompt to the freshly built flow, but reuseOAuthFlow could then answer from a cached flow for the same client — one built without prompt=login, e.g. by RequestJWTAuth. The user got the same silent authorization URL that produced the mismatch, with the flag already spent, so no later round asked either. Rule reuse out when the prompt is forced, while still cancelling the predecessor's wait. RequestJWTAuth also wrote the flow fields one by one, leaving the previous login's hint and accountPrompted behind for WaitSSOLogin to judge a later token against. Both sites now replace the whole record. * [client] Consume the forced account prompt after the retry forceAccountPrompt was never cleared, so a flow that outlived the retry it was armed for kept sending prompt=login on every later authorization request and re-authenticated the user each time. RequestAuthInfo now takes the flag as it builds the request. * [client] Cancel the caller context in the SSO login tests WaitSSOLogin parks a goroutine on the caller's context for the whole browser leg. The tests passed context.Background(), which never cancels, so each left one goroutine behind for the lifetime of the test binary. * [client] Cancel the wait displaced by an OAuth flow replacement Replacing the shared record with a whole struct value dropped the previous flow's waitCancel, so an SSO browser wait still parked on it lost its cancel: nothing could preempt it, and it could go on to run attemptLogin or mutate the record behind the new flow. Both replacement sites now take the displaced cancel over in the same critical section, via a shared replaceOAuthFlow, and invoke it after the unlock. * [client] Guard OAuth flow mutations by the flow that owns the wait * [client] Arm the account prompt only from the wait that owns the flow * [client] Drop the pending login flow when a login switches the profile SwitchProfile cancels the pending OAuth wait and clears the flow record, the account-prompt flag and the pending session extend, because they describe the previous profile's login. A Login or Up that carries a ProfileName switches the profile through switchProfileIfNeeded without that cleanup, so reuseOAuthFlow could hand the new profile the previous profile's flow: the same IdP client ID, the previous account's login_hint in the URL, and a record whose hint WaitSSOLogin would judge the new profile's token against. That either fails the login with a spurious account mismatch or lets a fresh peer register under the other account. switchProfileIfNeeded now reports whether it switched, and its callers run the same cleanup on a switch. A login on the same profile keeps the pending flow, so a second client can still join it. * [client] Clear the JWT cache when a login switches the profile The daemon keeps the user's JWT in a cache for SSH logins. When the user switches to another profile, this token belongs to the old profile, so it must not be used for the new one. SwitchProfile already cleared the cache. A login or up request can also switch the profile, but it did not clear the cache. After such a switch, SSH could still use the old profile's token, and a sign-in still running for the old profile could save its token into the new profile's cache. Clear the cache in dropPendingAuthFlows, so every profile switch does it. * [client] Bump the JWT cache generation with the config swap in Login RequestJWTAuth snapshots s.config and the cache generation under one s.mutex section and relies on the two flipping together. Login cleared the cache right after the profile switch but replaced s.config only later, after getConfig, so a JWT flow started in between carried the previous profile's config with the new generation, and its token landed in the cache the new profile then served. The swap now bumps the generation under the same lock. The early drop stays: it covers a login that fails after the switch, where a retry no longer sees a profile change. --- client/server/login_account_test.go | 121 ++++++++++++++++++++++++++++ client/server/server.go | 118 ++++++++++++++++----------- 2 files changed, 190 insertions(+), 49 deletions(-) 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 From a5834fdaab96834bde79303b247113d23ebb1380 Mon Sep 17 00:00:00 2001 From: Nicolas Frati Date: Fri, 9 Oct 2026 13:37:25 +0200 Subject: [PATCH 2/2] [management,signal] Make the Let's Encrypt challenge listener address configurable (#7706) * [management,signal] Make the Let's Encrypt challenge listener address configurable With Let's Encrypt enabled and --port set to something other than 443, signal and management also opened a separate challenge listener that was hard-coded to :443. Non-root deployments, such as the UBI images, could not start that listener. Add --letsencrypt-listen-address to both. It defaults to :443, so current behavior is unchanged. An empty value disables the separate listener for setups that forward public port 443 to --port, where the main TLS listener already answers TLS-ALPN-01 challenges. Signal now fails on startup when the challenge listener cannot bind, and exits non-zero when a server stops unexpectedly instead of exiting 0. A failure reported before the run loop waited was previously dropped. Management no longer opens a new :443 listener on shutdown just to close it. * [management,signal] Keep the challenge listener change additive Remove the Signal fail-fast changes from this PR. They change the behavior that existing installations see after an upgrade, so they move to a separate PR. If the challenge listener cannot bind, Signal now logs the error and continues. The main TLS listener still answers TLS-ALPN-01 challenges. Management keeps its previous behavior and stops with an error. The check for an empty address moves to the caller, so the function does not return a nil listener with a nil error. Also add assertion messages, guard a nil listener in a test cleanup, and add the flag to the Signal README. --- management/cmd/management.go | 1 + management/cmd/root.go | 2 + management/internals/server/server.go | 40 +++++++++++--- .../server/server_letsencrypt_test.go | 55 +++++++++++++++++++ signal/README.md | 1 + signal/cmd/run.go | 47 ++++++++++++---- signal/cmd/run_test.go | 52 ++++++++++++++++++ 7 files changed, 181 insertions(+), 17 deletions(-) create mode 100644 management/internals/server/server_letsencrypt_test.go create mode 100644 signal/cmd/run_test.go diff --git a/management/cmd/management.go b/management/cmd/management.go index fc6bd0a46..434ad8402 100644 --- a/management/cmd/management.go +++ b/management/cmd/management.go @@ -143,6 +143,7 @@ var ( MgmtPort: mgmtPort, MgmtMetricsPort: mgmtMetricsPort, DisableLegacyManagementPort: disableLegacyManagementPort, + LetsEncryptListenAddress: mgmtLetsencryptListen, DisableMetrics: disableMetrics, DisableGeoliteUpdate: disableGeoliteUpdate, UserDeleteFromIDPEnabled: userDeleteFromIDPEnabled, diff --git a/management/cmd/root.go b/management/cmd/root.go index ae03a09e8..466ee133f 100644 --- a/management/cmd/root.go +++ b/management/cmd/root.go @@ -29,6 +29,7 @@ var ( mgmtMetricsPort int disableLegacyManagementPort bool mgmtLetsencryptDomain string + mgmtLetsencryptListen string mgmtSingleAccModeDomain string certFile string certKey string @@ -70,6 +71,7 @@ func init() { mgmtCmd.Flags().StringVar(&mgmtDataDir, "datadir", defaultMgmtDataDir, "server data directory location") mgmtCmd.Flags().StringVar(&nbconfig.MgmtConfigPath, "config", defaultMgmtConfig, "Netbird config file location. Config params specified via command line (e.g. datadir) have a precedence over configuration from this file") mgmtCmd.Flags().StringVar(&mgmtLetsencryptDomain, "letsencrypt-domain", "", "a domain to issue Let's Encrypt certificate for. Enables TLS using Let's Encrypt. Will fetch and renew certificate, and run the server with TLS") + mgmtCmd.Flags().StringVar(&mgmtLetsencryptListen, "letsencrypt-listen-address", ":443", "address of the separate Let's Encrypt challenge listener, used when --port is not 443. Set it empty when public port 443 is forwarded to --port, which answers the challenges itself") mgmtCmd.Flags().StringVar(&mgmtSingleAccModeDomain, "single-account-mode-domain", defaultSingleAccModeDomain, "Enables single account mode. This means that all the users will be under the same account grouped by the specified domain. If the installation has more than one account, the property is ineffective. Enabled by default with the default domain "+defaultSingleAccModeDomain) mgmtCmd.Flags().BoolVar(&disableSingleAccMode, "disable-single-account-mode", false, "If set to true, disables single account mode. The --single-account-mode-domain property will be ignored and every new user will have a separate NetBird account.") mgmtCmd.Flags().StringVar(&certFile, "cert-file", "", "Location of your SSL certificate. Can be used when you have an existing certificate and don't want a new certificate be generated automatically. If letsencrypt-domain is specified this property has no effect") diff --git a/management/internals/server/server.go b/management/internals/server/server.go index 6d51745a7..bbf0f389d 100644 --- a/management/internals/server/server.go +++ b/management/internals/server/server.go @@ -68,6 +68,7 @@ type BaseServer struct { mgmtMetricsPort int mgmtPort int disableLegacyManagementPort bool + letsEncryptListenAddress string autoResolveDomains bool proxyAuthClose func() @@ -82,6 +83,8 @@ type BaseServer struct { tlsConfig *tls.Config certManager *autocert.Manager update *version.Update + // certListener serves Let's Encrypt challenges when mgmtPort is not 443. + certListener net.Listener errCh chan error wg sync.WaitGroup @@ -103,6 +106,9 @@ type Config struct { UserDeleteFromIDPEnabled bool AutoResolveDomains bool TLSConfig *tls.Config + // LetsEncryptListenAddress is the separate Let's Encrypt challenge listener + // used when MgmtPort is not 443. Empty disables it. + LetsEncryptListenAddress string } // NewServer initializes and configures a new Server instance @@ -117,6 +123,7 @@ func NewServer(cfg *Config) *BaseServer { userDeleteFromIDPEnabled: cfg.UserDeleteFromIDPEnabled, mgmtPort: cfg.MgmtPort, disableLegacyManagementPort: cfg.DisableLegacyManagementPort, + letsEncryptListenAddress: cfg.LetsEncryptListenAddress, mgmtMetricsPort: cfg.MgmtMetricsPort, autoResolveDomains: cfg.AutoResolveDomains, tlsConfig: cfg.TLSConfig, @@ -210,19 +217,18 @@ func (s *BaseServer) start(ctx context.Context) error { rootHandler := s.handlerFunc(srvCtx, s.GRPCServer(), s.APIHandler(), s.IDPHandler(), s.Metrics().GetMeter()) switch { case s.certManager != nil: - // a call to certManager.Listener() always creates a new listener so we do it once - cml := s.certManager.Listener() if s.mgmtPort == 443 { // CertManager, HTTP and gRPC API all on the same port rootHandler = s.certManager.HTTPHandler(rootHandler) - s.listener = cml + s.listener = s.certManager.Listener() } else { s.listener, err = tls.Listen("tcp", fmt.Sprintf(":%d", s.mgmtPort), s.certManager.TLSConfig()) if err != nil { return fmt.Errorf("failed creating TLS listener on port %d: %v", s.mgmtPort, err) } - log.WithContext(ctx).Infof("running HTTP server (LetsEncrypt challenge handler): %s", cml.Addr().String()) - s.serveHTTP(ctx, cml, s.certManager.HTTPHandler(nil)) + if err := s.serveLetsEncryptChallenges(ctx); err != nil { + return err + } } case s.tlsConfig != nil: s.listener, err = tls.Listen("tcp", fmt.Sprintf(":%d", s.mgmtPort), s.tlsConfig) @@ -309,8 +315,8 @@ func (s *BaseServer) Stop() error { if s.listener != nil { _ = s.listener.Close() } - if s.certManager != nil { - _ = s.certManager.Listener().Close() + if s.certListener != nil { + _ = s.certListener.Close() } s.GRPCServer().Stop() if s.proxyAuthClose != nil { @@ -416,6 +422,26 @@ func (s *BaseServer) serveGRPC(ctx context.Context, grpcServer *grpc.Server, por return listener, nil } +// serveLetsEncryptChallenges starts the separate Let's Encrypt challenge listener +// unless it is disabled. The main TLS listener uses the cert manager's TLS +// config, so it still answers TLS-ALPN-01 challenges when public port 443 is +// forwarded to it. +func (s *BaseServer) serveLetsEncryptChallenges(ctx context.Context) error { + if s.letsEncryptListenAddress == "" { + log.WithContext(ctx).Infof("LetsEncrypt challenge server disabled, challenges are answered on port %d", s.mgmtPort) + return nil + } + + cml, err := tls.Listen("tcp", s.letsEncryptListenAddress, s.certManager.TLSConfig()) + if err != nil { + return fmt.Errorf("create LetsEncrypt challenge listener on %s: %w", s.letsEncryptListenAddress, err) + } + s.certListener = cml + log.WithContext(ctx).Infof("running HTTP server (LetsEncrypt challenge handler): %s", cml.Addr().String()) + s.serveHTTP(ctx, cml, s.certManager.HTTPHandler(nil)) + return nil +} + func (s *BaseServer) serveHTTP(ctx context.Context, httpListener net.Listener, handler http.Handler) { s.wg.Add(1) go func() { diff --git a/management/internals/server/server_letsencrypt_test.go b/management/internals/server/server_letsencrypt_test.go new file mode 100644 index 000000000..0827d5e5d --- /dev/null +++ b/management/internals/server/server_letsencrypt_test.go @@ -0,0 +1,55 @@ +package server + +import ( + "context" + "net" + "testing" + "time" + + "github.com/stretchr/testify/require" + "golang.org/x/crypto/acme/autocert" + + nbconfig "github.com/netbirdio/netbird/management/internals/server/config" +) + +func newLetsEncryptTestServer(address string) *BaseServer { + srv := NewServer(&Config{NbConfig: &nbconfig.Config{}, MgmtPort: 8443, LetsEncryptListenAddress: address}) + srv.certManager = &autocert.Manager{} + return srv +} + +func TestServeLetsEncryptChallenges_Disabled(t *testing.T) { + srv := newLetsEncryptTestServer("") + + require.NoError(t, srv.serveLetsEncryptChallenges(context.Background())) + require.Nil(t, srv.certListener, "no challenge listener should be created when the address is empty") +} + +func TestServeLetsEncryptChallenges_CustomAddress(t *testing.T) { + srv := newLetsEncryptTestServer("127.0.0.1:0") + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(func() { + cancel() + if srv.certListener != nil { + _ = srv.certListener.Close() + } + srv.wg.Wait() + }) + + require.NoError(t, srv.serveLetsEncryptChallenges(ctx)) + require.NotNil(t, srv.certListener, "challenge listener should be created on the configured address") + + conn, err := net.DialTimeout("tcp", srv.certListener.Addr().String(), time.Second) + require.NoError(t, err) + require.NoError(t, conn.Close()) +} + +func TestServeLetsEncryptChallenges_BindFailure(t *testing.T) { + occupied, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + t.Cleanup(func() { _ = occupied.Close() }) + srv := newLetsEncryptTestServer(occupied.Addr().String()) + + require.Error(t, srv.serveLetsEncryptChallenges(context.Background())) + require.Nil(t, srv.certListener, "no challenge listener should be stored when the bind fails") +} diff --git a/signal/README.md b/signal/README.md index 0033eaf90..dd57252ad 100644 --- a/signal/README.md +++ b/signal/README.md @@ -16,6 +16,7 @@ Usage: Flags: -h, --help help for run --letsencrypt-domain string a domain to issue Let's Encrypt certificate for. Enables TLS using Let's Encrypt. Will fetch and renew certificate, and run the server with TLS + --letsencrypt-listen-address string address of the separate Let's Encrypt challenge listener, used when --port is not 443. Set it empty when public port 443 is forwarded to --port, which answers the challenges itself (default ":443") --port int Server port to listen on (e.g. 10000) (default 10000) --ssl-dir string server ssl directory location. *Required only for Let's Encrypt certificates. (default "/var/lib/netbird/") --cert-file string Location of your SSL certificate. Can be used when you have an existing certificate and don't want a new certificate be generated automatically. If letsencrypt-domain is specified this property has no effect diff --git a/signal/cmd/run.go b/signal/cmd/run.go index 42b7d2505..ffb596909 100644 --- a/signal/cmd/run.go +++ b/signal/cmd/run.go @@ -36,7 +36,12 @@ import ( "google.golang.org/grpc/keepalive" ) -const legacyGRPCPort = 10000 +const ( + legacyGRPCPort = 10000 + // defaultLetsencryptListenAddress is where Let's Encrypt connects for + // TLS-ALPN-01 challenges unless public port 443 is forwarded elsewhere. + defaultLetsencryptListenAddress = ":443" +) var ( signalPort int @@ -44,6 +49,7 @@ var ( signalLetsencryptDomain string signalLetsencryptEmail string signalLetsencryptDataDir string + signalLetsencryptListen string signalCertFile string signalCertKey string @@ -124,8 +130,18 @@ var ( grpcRootHandler := grpcHandlerFunc(grpcServer, metricsServer.Meter) - if certManager != nil { - startServerWithCertManager(certManager, grpcRootHandler) + var certListener net.Listener + switch { + case certManager == nil: + case signalPort != 443 && signalLetsencryptListen == "": + // The main TLS listener uses the cert manager's TLS config, so it still + // answers TLS-ALPN-01 challenges when public port 443 is forwarded to it. + log.Infof("LetsEncrypt challenge server disabled, challenges are answered on port %d", signalPort) + default: + certListener, err = startServerWithCertManager(certManager, grpcRootHandler) + if err != nil { + log.Errorf("LetsEncrypt challenge server not started: %v", err) + } } var compatListener net.Listener @@ -169,6 +185,10 @@ var ( SetupCloseHandler() <-stopCh + if certListener != nil { + _ = certListener.Close() + log.Infof("stopped LetsEncrypt challenge server") + } if grpcListener != nil { _ = grpcListener.Close() log.Infof("stopped gRPC server") @@ -245,18 +265,24 @@ func getTLSConfigurations() ([]grpc.ServerOption, *autocert.Manager, *tls.Config return []grpc.ServerOption{grpc.Creds(transportCredentials)}, certManager, tlsConfig, err } -func startServerWithCertManager(certManager *autocert.Manager, grpcRootHandler http.Handler) { - // a call to certManager.Listener() always creates a new listener so we do it once - httpListener := certManager.Listener() +func startServerWithCertManager(certManager *autocert.Manager, grpcRootHandler http.Handler) (net.Listener, error) { if signalPort == 443 { + // a call to certManager.Listener() always creates a new listener so we do it once + httpListener := certManager.Listener() // running gRPC and HTTP cert manager on the same port serveHTTP(httpListener, certManager.HTTPHandler(grpcRootHandler)) log.Infof("running HTTP server (LetsEncrypt challenge handler) and gRPC server on the same port: %s", httpListener.Addr().String()) - } else { - // Start the HTTP cert manager server separately - serveHTTP(httpListener, certManager.HTTPHandler(nil)) - log.Infof("running HTTP server (LetsEncrypt challenge handler): %s", httpListener.Addr().String()) + return httpListener, nil } + + httpListener, err := tls.Listen("tcp", signalLetsencryptListen, certManager.TLSConfig()) + if err != nil { + return nil, fmt.Errorf("create LetsEncrypt challenge listener on %s: %w", signalLetsencryptListen, err) + } + // Start the HTTP cert manager server separately + serveHTTP(httpListener, certManager.HTTPHandler(nil)) + log.Infof("running HTTP server (LetsEncrypt challenge handler): %s", httpListener.Addr().String()) + return httpListener, nil } func grpcHandlerFunc(grpcServer *grpc.Server, meter metric.Meter) http.Handler { @@ -334,6 +360,7 @@ func init() { runCmd.PersistentFlags().StringVar(&signalLetsencryptDataDir, "letsencrypt-data-dir", "", "a directory to store Let's Encrypt data. Required if Let's Encrypt is enabled.") runCmd.PersistentFlags().StringVar(&signalLetsencryptDataDir, "ssl-dir", "", "server ssl directory location. *Required only for Let's Encrypt certificates. Deprecated: use --letsencrypt-data-dir") runCmd.PersistentFlags().StringVar(&signalLetsencryptDomain, "letsencrypt-domain", "", "a domain to issue Let's Encrypt certificate for. Enables TLS using Let's Encrypt. Will fetch and renew certificate, and run the server with TLS") + runCmd.PersistentFlags().StringVar(&signalLetsencryptListen, "letsencrypt-listen-address", defaultLetsencryptListenAddress, "address of the separate Let's Encrypt challenge listener, used when --port is not 443. Set it empty when public port 443 is forwarded to --port, which answers the challenges itself") runCmd.PersistentFlags().StringVar(&signalLetsencryptEmail, "letsencrypt-email", "", "email address to use for Let's Encrypt certificate registration") runCmd.PersistentFlags().StringVar(&signalCertFile, "cert-file", "", "Location of your SSL certificate. Can be used when you have an existing certificate and don't want a new certificate be generated automatically. If letsencrypt-domain is specified this property has no effect") runCmd.PersistentFlags().StringVar(&signalCertKey, "cert-key", "", "Location of your SSL certificate private key. Can be used when you have an existing certificate and don't want a new certificate be generated automatically. If letsencrypt-domain is specified this property has no effect") diff --git a/signal/cmd/run_test.go b/signal/cmd/run_test.go new file mode 100644 index 000000000..e0a07e923 --- /dev/null +++ b/signal/cmd/run_test.go @@ -0,0 +1,52 @@ +package cmd + +import ( + "net" + "net/http" + "testing" + "time" + + "github.com/stretchr/testify/require" + "golang.org/x/crypto/acme" + "golang.org/x/crypto/acme/autocert" +) + +func setLetsencryptListen(t *testing.T, port int, address string) { + t.Helper() + oldPort, oldAddress := signalPort, signalLetsencryptListen + signalPort, signalLetsencryptListen = port, address + t.Cleanup(func() { + signalPort, signalLetsencryptListen = oldPort, oldAddress + }) +} + +func TestStartServerWithCertManager_CustomAddress(t *testing.T) { + setLetsencryptListen(t, 10000, "127.0.0.1:0") + + listener, err := startServerWithCertManager(&autocert.Manager{}, http.NotFoundHandler()) + require.NoError(t, err) + require.NotNil(t, listener, "challenge listener should be created on the configured address") + t.Cleanup(func() { _ = listener.Close() }) + + conn, err := net.DialTimeout("tcp", listener.Addr().String(), time.Second) + require.NoError(t, err) + require.NoError(t, conn.Close()) +} + +func TestStartServerWithCertManager_BindFailure(t *testing.T) { + occupied, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + t.Cleanup(func() { _ = occupied.Close() }) + setLetsencryptListen(t, 10000, occupied.Addr().String()) + + listener, err := startServerWithCertManager(&autocert.Manager{}, http.NotFoundHandler()) + require.Error(t, err) + require.Nil(t, listener, "no listener should be returned when the bind fails") +} + +func TestCertManagerTLSConfigAnswersTLSALPN01(t *testing.T) { + // Disabling the separate listener relies on the main listener answering + // TLS-ALPN-01 challenges through the cert manager's TLS config. + cfg := (&autocert.Manager{}).TLSConfig() + require.Contains(t, cfg.NextProtos, acme.ALPNProto, "cert manager TLS config should offer the ACME TLS-ALPN protocol") +}