diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go index 244bda0b5..b9d490dd9 100644 --- a/management/server/store/sql_store.go +++ b/management/server/store/sql_store.go @@ -5929,7 +5929,7 @@ func (s *SqlStore) getClusterCapability(ctx context.Context, clusterAddr, column AnyTrue bool } - err := s.db.WithContext(ctx). + err := s.db. Model(&proxy.Proxy{}). Select("COUNT(CASE WHEN "+column+" IS NOT NULL THEN 1 END) > 0 AS has_capability, "+ "COALESCE(MAX(CASE WHEN "+column+" = true THEN 1 ELSE 0 END), 0) = 1 AS any_true"). diff --git a/proxy/cmd/proxy/cmd/root.go b/proxy/cmd/proxy/cmd/root.go index 9af26f09f..5970886da 100644 --- a/proxy/cmd/proxy/cmd/root.go +++ b/proxy/cmd/proxy/cmd/root.go @@ -106,11 +106,7 @@ func init() { rootCmd.Flags().StringVar(&preSharedKey, "preshared-key", envStringOrDefault("NB_PROXY_PRESHARED_KEY", ""), "Define a pre-shared key for the tunnel between proxy and peers") rootCmd.Flags().BoolVar(&supportsCustomPorts, "supports-custom-ports", envBoolOrDefault("NB_PROXY_SUPPORTS_CUSTOM_PORTS", true), "Whether the proxy can bind arbitrary ports for UDP/TCP passthrough") rootCmd.Flags().BoolVar(&requireSubdomain, "require-subdomain", envBoolOrDefault("NB_PROXY_REQUIRE_SUBDOMAIN", false), "Require a subdomain label in front of the cluster domain") - // --private is internal: set by the embedded `netbird proxy` subcommand - // via NB_PROXY_PRIVATE so management can distinguish per-peer / private - // clusters from centralised ones. Hidden so the standalone CLI doesn't - // surface it as an operator-facing toggle. - rootCmd.Flags().BoolVar(&private, "private", envBoolOrDefault("NB_PROXY_PRIVATE", false), "Mark this proxy as embedded/private (internal flag)") + rootCmd.Flags().BoolVar(&private, "private", envBoolOrDefault("NB_PROXY_PRIVATE", false), "Enable private services accessible with NetBird-Only authentication mode.") _ = rootCmd.Flags().MarkHidden("private") rootCmd.Flags().DurationVar(&maxDialTimeout, "max-dial-timeout", envDurationOrDefault("NB_PROXY_MAX_DIAL_TIMEOUT", 0), "Cap per-service backend dial timeout (0 = no cap)") rootCmd.Flags().DurationVar(&maxSessionIdleTimeout, "max-session-idle-timeout", envDurationOrDefault("NB_PROXY_MAX_SESSION_IDLE_TIMEOUT", 0), "Cap per-service session idle timeout (0 = no cap)") diff --git a/proxy/internal/auth/middleware.go b/proxy/internal/auth/middleware.go index 1b93962b6..ee8cc6541 100644 --- a/proxy/internal/auth/middleware.go +++ b/proxy/internal/auth/middleware.go @@ -311,12 +311,6 @@ func (mw *Middleware) forwardWithSessionCookie(w http.ResponseWriter, r *http.Re if err != nil { return false } - if requestIsPlainHTTP(r) { - mw.logger.WithFields(log.Fields{ - "host": host, - "remote": r.RemoteAddr, - }).Warn("session cookie on plain HTTP path; cookie auth requires TLS — use port 443") - } userID, email, method, groups, groupNames, err := auth.ValidateSessionJWT(cookie.Value, host, config.SessionPublicKey) if err != nil { return false @@ -696,8 +690,6 @@ func (mw *Middleware) validateSessionToken(ctx context.Context, host, token stri UserEmail: resp.GetUserEmail(), Valid: false, DeniedReason: resp.DeniedReason, - Groups: resp.GetPeerGroupIds(), - GroupNames: resp.GetPeerGroupNames(), }, nil } return &validationResult{ diff --git a/proxy/lifecycle_test.go b/proxy/lifecycle_test.go deleted file mode 100644 index a35750df6..000000000 --- a/proxy/lifecycle_test.go +++ /dev/null @@ -1,134 +0,0 @@ -package proxy - -import ( - "context" - "errors" - "io" - "testing" - "time" - - log "github.com/sirupsen/logrus" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -// quietLifecycleLogger keeps lifecycle tests from spamming the test output. -func quietLifecycleLogger() *log.Logger { - l := log.New() - l.SetOutput(io.Discard) - l.SetLevel(log.PanicLevel) - return l -} - -func TestNewIsPureConstructor(t *testing.T) { - cfg := Config{ - ListenAddr: ":0", - ID: "test-id", - Logger: quietLifecycleLogger(), - Version: "test", - ManagementAddress: "https://example.invalid", - HealthAddr: "", - ForwardedProto: "auto", - } - - srv := New(cfg) - require.NotNil(t, srv, "New must return a non-nil Server") - - assert.Equal(t, ":0", srv.ListenAddr, "ListenAddr should round-trip") - assert.Equal(t, "test-id", srv.ID, "ID should round-trip") - assert.Equal(t, "test", srv.Version, "Version should round-trip") - assert.Equal(t, "https://example.invalid", srv.ManagementAddress, "ManagementAddress should round-trip") - assert.Equal(t, "auto", srv.ForwardedProto, "ForwardedProto should round-trip") - - // Pure constructor: no goroutines, no listener bind, no management dial. - assert.False(t, srv.started, "Server must be marked unstarted before Start") - assert.Nil(t, srv.mgmtClient, "mgmt client must not be created in New") - assert.Nil(t, srv.netbird, "netbird client must not be created in New") - assert.Nil(t, srv.https, "https server must not be created in New") - assert.Nil(t, srv.healthServer, "health server must not be created in New") - assert.Nil(t, srv.runCancel, "runCancel must be nil before Start") - assert.Nil(t, srv.runErrCh, "runErrCh must be nil before Start") -} - -func TestStopBeforeStartIsNoOp(t *testing.T) { - srv := New(Config{Logger: quietLifecycleLogger()}) - - ctx, cancel := context.WithTimeout(context.Background(), time.Second) - defer cancel() - - err := srv.Stop(ctx) - assert.NoError(t, err, "Stop on an unstarted server must succeed without error") - - err = srv.Stop(ctx) - assert.NoError(t, err, "Stop must remain idempotent across repeated calls") -} - -func TestStartFailsWithoutManagement(t *testing.T) { - srv := New(Config{ - Logger: quietLifecycleLogger(), - ListenAddr: "127.0.0.1:0", - ManagementAddress: "://broken-url", - }) - - ctx, cancel := context.WithTimeout(context.Background(), time.Second) - defer cancel() - - err := srv.Start(ctx) - require.Error(t, err, "Start must surface management dial failures") - - assert.True(t, srv.started, "started flag is set before any dial attempt so a second Start fails fast") - - err = srv.Start(ctx) - require.Error(t, err, "second Start must reject") - assert.Contains(t, err.Error(), "already started", "error must explain why the call was rejected") -} - -func TestStopIsIdempotent(t *testing.T) { - srv := &Server{ - Logger: quietLifecycleLogger(), - started: true, - runErrCh: make(chan struct{}), - runCancel: func() {}, - } - srv.recordRunErr(errors.New("synthetic")) - - ctx, cancel := context.WithTimeout(context.Background(), time.Second) - defer cancel() - - err := srv.Stop(ctx) - require.Error(t, err, "Stop must surface the recorded background error") - assert.Contains(t, err.Error(), "synthetic", "error must round-trip recordRunErr's value") - - err = srv.Stop(ctx) - require.Error(t, err, "second Stop must still report the same error") - assert.Contains(t, err.Error(), "synthetic", "idempotent Stop must return the cached error") -} - -func TestRecordRunErrPreservesFirstFailure(t *testing.T) { - srv := &Server{ - Logger: quietLifecycleLogger(), - runErrCh: make(chan struct{}), - } - - srv.recordRunErr(errors.New("first")) - srv.recordRunErr(errors.New("second")) - - require.Error(t, srv.runErr, "first failure must be retained") - assert.Contains(t, srv.runErr.Error(), "first", "second call must not overwrite the cached error") - - select { - case <-srv.runErrCh: - default: - t.Fatal("recordRunErr must close runErrCh so waitAndStop unblocks") - } -} - -func TestStopSkipsShutdownWhenNeverStarted(t *testing.T) { - srv := New(Config{Logger: quietLifecycleLogger()}) - - ctx, cancel := context.WithCancel(context.Background()) - cancel() - - err := srv.Stop(ctx) - assert.NoError(t, err, "Stop on an unstarted server should not block on the cancelled ctx") -} diff --git a/proxy/server_test.go b/proxy/server_test.go index b4fb4f8ba..549ac65f0 100644 --- a/proxy/server_test.go +++ b/proxy/server_test.go @@ -1,9 +1,15 @@ package proxy import ( + "context" + "errors" + "io" "testing" + "time" + log "github.com/sirupsen/logrus" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestDebugEndpointDisabledByDefault(t *testing.T) { @@ -46,3 +52,94 @@ func TestDebugEndpointAddr(t *testing.T) { }) } } + +// quietLifecycleLogger keeps lifecycle tests from spamming the test output. +func quietLifecycleLogger() *log.Logger { + l := log.New() + l.SetOutput(io.Discard) + l.SetLevel(log.PanicLevel) + return l +} + +func TestStopBeforeStartIsNoOp(t *testing.T) { + srv := New(Config{Logger: quietLifecycleLogger()}) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + + err := srv.Stop(ctx) + assert.NoError(t, err, "Stop on an unstarted server must succeed without error") + + err = srv.Stop(ctx) + assert.NoError(t, err, "Stop must remain idempotent across repeated calls") +} + +func TestStartFailsWithoutManagement(t *testing.T) { + srv := New(Config{ + Logger: quietLifecycleLogger(), + ListenAddr: "127.0.0.1:0", + ManagementAddress: "://broken-url", + }) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + + err := srv.Start(ctx) + require.Error(t, err, "Start must surface management dial failures") + + assert.True(t, srv.started, "started flag is set before any dial attempt so a second Start fails fast") + + err = srv.Start(ctx) + require.Error(t, err, "second Start must reject") + assert.Contains(t, err.Error(), "already started", "error must explain why the call was rejected") +} + +func TestStopIsIdempotent(t *testing.T) { + srv := &Server{ + Logger: quietLifecycleLogger(), + started: true, + runErrCh: make(chan struct{}), + runCancel: func() {}, + } + srv.recordRunErr(errors.New("synthetic")) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + + err := srv.Stop(ctx) + require.Error(t, err, "Stop must surface the recorded background error") + assert.Contains(t, err.Error(), "synthetic", "error must round-trip recordRunErr's value") + + err = srv.Stop(ctx) + require.Error(t, err, "second Stop must still report the same error") + assert.Contains(t, err.Error(), "synthetic", "idempotent Stop must return the cached error") +} + +func TestRecordRunErrPreservesFirstFailure(t *testing.T) { + srv := &Server{ + Logger: quietLifecycleLogger(), + runErrCh: make(chan struct{}), + } + + srv.recordRunErr(errors.New("first")) + srv.recordRunErr(errors.New("second")) + + require.Error(t, srv.runErr, "first failure must be retained") + assert.Contains(t, srv.runErr.Error(), "first", "second call must not overwrite the cached error") + + select { + case <-srv.runErrCh: + default: + t.Fatal("recordRunErr must close runErrCh so waitAndStop unblocks") + } +} + +func TestStopSkipsShutdownWhenNeverStarted(t *testing.T) { + srv := New(Config{Logger: quietLifecycleLogger()}) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + err := srv.Stop(ctx) + assert.NoError(t, err, "Stop on an unstarted server should not block on the cancelled ctx") +}