mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-10 07:29:06 +02:00
Merge remote-tracking branch 'origin/main' into jnfrati/ubi-signal
This commit is contained in:
@@ -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
|
||||
|
||||
+69
-49
@@ -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
|
||||
|
||||
@@ -143,6 +143,7 @@ var (
|
||||
MgmtPort: mgmtPort,
|
||||
MgmtMetricsPort: mgmtMetricsPort,
|
||||
DisableLegacyManagementPort: disableLegacyManagementPort,
|
||||
LetsEncryptListenAddress: mgmtLetsencryptListen,
|
||||
DisableMetrics: disableMetrics,
|
||||
DisableGeoliteUpdate: disableGeoliteUpdate,
|
||||
UserDeleteFromIDPEnabled: userDeleteFromIDPEnabled,
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
|
||||
+37
-10
@@ -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")
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
Reference in New Issue
Block a user