diff --git a/management/internals/modules/reverseproxy/proxy/manager.go b/management/internals/modules/reverseproxy/proxy/manager.go index c0b8435ec..9350ad9b9 100644 --- a/management/internals/modules/reverseproxy/proxy/manager.go +++ b/management/internals/modules/reverseproxy/proxy/manager.go @@ -20,6 +20,7 @@ type Manager interface { ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool + ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool CleanupStale(ctx context.Context, inactivityDuration time.Duration) error GetAccountProxy(ctx context.Context, accountID string) (*Proxy, error) diff --git a/management/internals/modules/reverseproxy/proxy/manager/manager.go b/management/internals/modules/reverseproxy/proxy/manager/manager.go index 7ddb66eec..5a95ea94a 100644 --- a/management/internals/modules/reverseproxy/proxy/manager/manager.go +++ b/management/internals/modules/reverseproxy/proxy/manager/manager.go @@ -23,6 +23,7 @@ type store interface { GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool + GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error) CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error GetProxyByAccountID(ctx context.Context, accountID string) (*proxy.Proxy, error) @@ -149,6 +150,11 @@ func (m Manager) ClusterSupportsPrivate(ctx context.Context, clusterAddr string) return m.store.GetClusterSupportsPrivate(ctx, clusterAddr) } +// ClusterAllProxiesPrivate reports whether every active proxy claims the private capability (nil = unreported). +func (m Manager) ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool { + return m.store.GetClusterAllProxiesPrivate(ctx, clusterAddr) +} + // ClusterSupportsSessionCode reports whether all active proxies support session codes. func (m Manager) ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool { versions, err := m.store.GetActiveProxyVersions(ctx, clusterAddr) diff --git a/management/internals/modules/reverseproxy/proxy/manager/manager_test.go b/management/internals/modules/reverseproxy/proxy/manager/manager_test.go index 66ddb95bd..56806613a 100644 --- a/management/internals/modules/reverseproxy/proxy/manager/manager_test.go +++ b/management/internals/modules/reverseproxy/proxy/manager/manager_test.go @@ -105,6 +105,9 @@ func (m *mockStore) GetClusterSupportsCrowdSec(_ context.Context, _ string) *boo func (m *mockStore) GetClusterSupportsPrivate(_ context.Context, _ string) *bool { return nil } +func (m *mockStore) GetClusterAllProxiesPrivate(_ context.Context, _ string) *bool { + return nil +} func (m *mockStore) GetActiveProxyVersions(ctx context.Context, clusterAddress string) ([]string, error) { if m.getActiveProxyVersionsFunc != nil { return m.getActiveProxyVersionsFunc(ctx, clusterAddress) diff --git a/management/internals/modules/reverseproxy/proxy/manager_mock.go b/management/internals/modules/reverseproxy/proxy/manager_mock.go index d6f7197d7..5f3404096 100644 --- a/management/internals/modules/reverseproxy/proxy/manager_mock.go +++ b/management/internals/modules/reverseproxy/proxy/manager_mock.go @@ -56,6 +56,20 @@ func (mr *MockManagerMockRecorder) CleanupStale(ctx, inactivityDuration any) *go return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CleanupStale", reflect.TypeOf((*MockManager)(nil).CleanupStale), ctx, inactivityDuration) } +// ClusterAllProxiesPrivate mocks base method. +func (m *MockManager) ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ClusterAllProxiesPrivate", ctx, clusterAddr) + ret0, _ := ret[0].(*bool) + return ret0 +} + +// ClusterAllProxiesPrivate indicates an expected call of ClusterAllProxiesPrivate. +func (mr *MockManagerMockRecorder) ClusterAllProxiesPrivate(ctx, clusterAddr any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterAllProxiesPrivate", reflect.TypeOf((*MockManager)(nil).ClusterAllProxiesPrivate), ctx, clusterAddr) +} + // ClusterRequireSubdomain mocks base method. func (m *MockManager) ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool { m.ctrl.T.Helper() diff --git a/management/internals/modules/reverseproxy/service/manager/manager.go b/management/internals/modules/reverseproxy/service/manager/manager.go index 62897c9ae..900b7759f 100644 --- a/management/internals/modules/reverseproxy/service/manager/manager.go +++ b/management/internals/modules/reverseproxy/service/manager/manager.go @@ -84,6 +84,7 @@ type CapabilityProvider interface { ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool + ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool } type Manager struct { @@ -332,6 +333,10 @@ func (m *Manager) persistNewService(ctx context.Context, accountID string, svc * return err } + if err := m.validatePrivateClusterTargets(ctx, svc.Targets, svc.ProxyCluster); err != nil { + return err + } + return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil { return err @@ -369,6 +374,43 @@ func (m *Manager) clusterCustomPorts(ctx context.Context, svc *service.Service) return m.capabilities.ClusterSupportsCustomPorts(ctx, svc.ProxyCluster) } +// validatePrivateClusterTargets rejects cluster and direct upstream targets unless +// every active proxy in the service's cluster reports the private capability. The +// mapping reaches all proxies in the cluster, so one non-private proxy would serve +// these targets too. An unreported capability is treated as unsupported. Must be +// called outside a transaction, like clusterCustomPorts. +func (m *Manager) validatePrivateClusterTargets(ctx context.Context, targets []*service.Target, cluster string) error { + target := firstPrivateClusterTarget(targets) + if target == nil { + return nil + } + + if private := m.capabilities.ClusterAllProxiesPrivate(ctx, cluster); private != nil && *private { + return nil + } + + if target.TargetType == service.TargetTypeCluster { + return status.Errorf(status.InvalidArgument, + "target_type %q requires a proxy cluster with private mode enabled, cluster %s does not support it", + service.TargetTypeCluster, cluster) + } + return status.Errorf(status.InvalidArgument, + "direct_upstream requires a proxy cluster with private mode enabled, cluster %s does not support it", cluster) +} + +// firstPrivateClusterTarget returns the first target that only a private cluster may serve. +func firstPrivateClusterTarget(targets []*service.Target) *service.Target { + for _, target := range targets { + if target == nil { + continue + } + if target.TargetType == service.TargetTypeCluster || target.Options.DirectUpstream { + return target + } + } + return nil +} + // ensureL4Port auto-assigns a listen port when needed and validates cluster support. // customPorts must be pre-computed via clusterCustomPorts before entering a transaction. func (m *Manager) ensureL4Port(ctx context.Context, tx store.Store, svc *service.Service, customPorts *bool, serviceUpdate bool) error { @@ -464,6 +506,10 @@ func (m *Manager) persistNewEphemeralService(ctx context.Context, accountID, pee return err } + if err := m.validatePrivateClusterTargets(ctx, svc.Targets, svc.ProxyCluster); err != nil { + return err + } + return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil { return err @@ -584,6 +630,10 @@ func (m *Manager) persistServiceUpdate(ctx context.Context, accountID string, se return nil, err } + if err := m.validatePrivateClusterTargets(ctx, service.Targets, effectiveCluster); err != nil { + return nil, err + } + // Validate subdomain requirement *before* the transaction: the underlying // capability lookup talks to the main DB pool, and SQLite's single-connection // pool would self-deadlock if this ran while the tx already held the only diff --git a/management/internals/modules/reverseproxy/service/manager/private_cluster_test.go b/management/internals/modules/reverseproxy/service/manager/private_cluster_test.go new file mode 100644 index 000000000..1f507294e --- /dev/null +++ b/management/internals/modules/reverseproxy/service/manager/private_cluster_test.go @@ -0,0 +1,218 @@ +package manager + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel/metric/noop" + "go.uber.org/mock/gomock" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy" + proxymanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy/manager" + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/shared/management/status" +) + +// setupPrivateClusterTest wires the real proxy manager as the capability +// provider and connects one proxy to testCluster reporting the given private +// capability. A nil private connects no proxy, so the capability is unreported. +func setupPrivateClusterTest(t *testing.T, private *bool) (*Manager, store.Store) { + t.Helper() + + mgr, testStore := setupIntegrationTest(t) + + proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter("")) + require.NoError(t, err) + mgr.capabilities = proxyMgr + + if private != nil { + connectTestProxy(t, proxyMgr, "proxy-1", &proxy.Capabilities{Private: private}) + } + + return mgr, testStore +} + +func connectTestProxy(t *testing.T, proxyMgr *proxymanager.Manager, proxyID string, caps *proxy.Capabilities) { + t.Helper() + _, err := proxyMgr.Connect(context.Background(), proxyID, "session-"+proxyID, testCluster, "127.0.0.1", "", nil, caps) + require.NoError(t, err) +} + +func clusterTarget() *rpservice.Target { + return &rpservice.Target{ + TargetId: testCluster, + TargetType: rpservice.TargetTypeCluster, + Host: "backend.lan", + Port: 8080, + Protocol: "http", + Enabled: true, + Options: rpservice.TargetOptions{DirectUpstream: true}, + } +} + +func directUpstreamPeerTarget() *rpservice.Target { + return &rpservice.Target{ + TargetId: testPeerID, + TargetType: rpservice.TargetTypePeer, + Host: "backend.lan", + Port: 8080, + Protocol: "http", + Enabled: true, + Options: rpservice.TargetOptions{DirectUpstream: true}, + } +} + +func TestCreateService_PrivateClusterTargets(t *testing.T) { + tests := []struct { + name string + private *bool + target *rpservice.Target + wantErr string + }{ + {name: "cluster target on private cluster", private: boolPtr(true), target: clusterTarget()}, + {name: "direct upstream on private cluster", private: boolPtr(true), target: directUpstreamPeerTarget()}, + {name: "cluster target on non-private cluster", private: boolPtr(false), target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`}, + {name: "direct upstream on non-private cluster", private: boolPtr(false), target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"}, + {name: "cluster target with unreported capability", private: nil, target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`}, + {name: "direct upstream with unreported capability", private: nil, target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupPrivateClusterTest(t, tc.private) + + svc := newTestService("app.test.netbird.io") + svc.Targets = []*rpservice.Target{tc.target} + + _, err := mgr.CreateService(ctx, testAccountID, testUserID, svc) + + services, listErr := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID) + require.NoError(t, listErr) + + if tc.wantErr == "" { + require.NoError(t, err) + assert.Len(t, services, 1, "the service should be persisted") + return + } + + require.Error(t, err) + assert.Contains(t, err.Error(), tc.wantErr) + sErr, ok := status.FromError(err) + require.True(t, ok, "the caller must receive a typed error") + assert.Equal(t, status.InvalidArgument, sErr.Type(), "the rejection should be an invalid argument") + assert.Empty(t, services, "a rejected service must not be persisted") + }) + } +} + +// A cluster where only some proxies run in private mode must not accept these +// targets: the mapping is delivered to every proxy in the cluster, so the +// non-private ones would serve the target from their host network as well. +func TestCreateService_MixedClusterRejectsPrivateTargets(t *testing.T) { + tests := []struct { + name string + secondCaps *proxy.Capabilities + }{ + {name: "second proxy reports not private", secondCaps: &proxy.Capabilities{Private: boolPtr(false)}}, + {name: "second proxy predates capability reporting", secondCaps: nil}, + } + + for _, tc := range tests { + for _, target := range []*rpservice.Target{clusterTarget(), directUpstreamPeerTarget()} { + t.Run(tc.name+"/"+string(target.TargetType), func(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupPrivateClusterTest(t, boolPtr(true)) + connectTestProxy(t, mgr.capabilities.(*proxymanager.Manager), "proxy-2", tc.secondCaps) + + svc := newTestService("app.test.netbird.io") + svc.Targets = []*rpservice.Target{target} + + _, err := mgr.CreateService(ctx, testAccountID, testUserID, svc) + require.Error(t, err, "a cluster with a non-private proxy must not accept the target") + assert.Contains(t, err.Error(), "requires a proxy cluster with private mode enabled") + + services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID) + require.NoError(t, err) + assert.Empty(t, services, "a rejected service must not be persisted") + }) + } + } +} + +func TestCreateService_RegularTargetIgnoresPrivateCapability(t *testing.T) { + ctx := context.Background() + mgr, _ := setupPrivateClusterTest(t, boolPtr(false)) + + _, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io")) + require.NoError(t, err, "a peer target without direct upstream must not need a private cluster") +} + +func TestUpdateService_PrivateClusterTargets(t *testing.T) { + tests := []struct { + name string + target *rpservice.Target + wantErr string + }{ + {name: "switch to cluster target", target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`}, + {name: "enable direct upstream", target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupPrivateClusterTest(t, boolPtr(false)) + + created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io")) + require.NoError(t, err) + + updated := newTestService("app.test.netbird.io") + updated.ID = created.ID + updated.AccountID = testAccountID + updated.Targets = []*rpservice.Target{tc.target} + + _, err = mgr.UpdateService(ctx, testAccountID, testUserID, updated) + require.Error(t, err) + assert.Contains(t, err.Error(), tc.wantErr) + + stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID) + require.NoError(t, err) + require.Len(t, stored.Targets, 1) + assert.Equal(t, rpservice.TargetTypePeer, stored.Targets[0].TargetType, "the stored target must be unchanged") + assert.False(t, stored.Targets[0].Options.DirectUpstream, "the stored target must keep direct upstream disabled") + }) + } +} + +func TestUpdateService_PrivateClusterAllowsClusterTarget(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupPrivateClusterTest(t, boolPtr(true)) + + created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io")) + require.NoError(t, err) + + updated := newTestService("app.test.netbird.io") + updated.ID = created.ID + updated.AccountID = testAccountID + updated.Targets = []*rpservice.Target{clusterTarget()} + + _, err = mgr.UpdateService(ctx, testAccountID, testUserID, updated) + require.NoError(t, err) + + stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID) + require.NoError(t, err) + require.Len(t, stored.Targets, 1) + assert.Equal(t, rpservice.TargetTypeCluster, stored.Targets[0].TargetType, "the cluster target should be stored") +} + +func TestValidatePrivateClusterTargets_NoLookupWithoutPrivateTargets(t *testing.T) { + ctrl := gomock.NewController(t) + // No ClusterAllProxiesPrivate expectation: a lookup would fail the test. + mgr := &Manager{capabilities: proxy.NewMockManager(ctrl)} + + targets := []*rpservice.Target{{TargetId: testPeerID, TargetType: rpservice.TargetTypePeer}} + require.NoError(t, mgr.validatePrivateClusterTargets(context.Background(), targets, testCluster)) +} diff --git a/management/server/store/sql_store_proxy.go b/management/server/store/sql_store_proxy.go index 58fa86468..bdccd282c 100644 --- a/management/server/store/sql_store_proxy.go +++ b/management/server/store/sql_store_proxy.go @@ -358,6 +358,14 @@ func (s *SqlStore) GetClusterSupportsPrivate(ctx context.Context, clusterAddr st return s.getClusterCapability(ctx, clusterAddr, "private") } +// GetClusterAllProxiesPrivate reports whether every active proxy in the cluster +// has the private capability. Returns nil when no proxy reported the capability. +// Use it where any proxy in the cluster may serve the result, since a single +// non-private proxy would serve it without the private guarantees. +func (s *SqlStore) GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool { + return s.getClusterUnanimousCapability(ctx, clusterAddr, "private") +} + // GetClusterSupportsCrowdSec returns whether all active proxies in the cluster // have CrowdSec configured. Returns nil when no proxy reported the capability. // Unlike other capabilities that use ANY-true (for rolling upgrades), CrowdSec diff --git a/management/server/store/store.go b/management/server/store/store.go index 177c2a47c..01aaf4892 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -336,6 +336,7 @@ type Store interface { GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool + GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error) CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error GetAllProxies(ctx context.Context) ([]*proxy.Proxy, error) diff --git a/management/server/store/store_mock.go b/management/server/store/store_mock.go index 53e35b866..956cac4b8 100644 --- a/management/server/store/store_mock.go +++ b/management/server/store/store_mock.go @@ -1870,6 +1870,20 @@ func (mr *MockStoreMockRecorder) GetAnyAccountID(ctx any) *gomock.Call { return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAnyAccountID", reflect.TypeOf((*MockStore)(nil).GetAnyAccountID), ctx) } +// GetClusterAllProxiesPrivate mocks base method. +func (m *MockStore) GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetClusterAllProxiesPrivate", ctx, clusterAddr) + ret0, _ := ret[0].(*bool) + return ret0 +} + +// GetClusterAllProxiesPrivate indicates an expected call of GetClusterAllProxiesPrivate. +func (mr *MockStoreMockRecorder) GetClusterAllProxiesPrivate(ctx, clusterAddr any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetClusterAllProxiesPrivate", reflect.TypeOf((*MockStore)(nil).GetClusterAllProxiesPrivate), ctx, clusterAddr) +} + // GetClusterRequireSubdomain mocks base method. func (m *MockStore) GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool { m.ctrl.T.Helper() diff --git a/proxy/management_integration_test.go b/proxy/management_integration_test.go index 000d8ce72..03a9855de 100644 --- a/proxy/management_integration_test.go +++ b/proxy/management_integration_test.go @@ -246,6 +246,10 @@ func (m *testProxyManager) ClusterSupportsPrivate(_ context.Context, _ string) * return nil } +func (m *testProxyManager) ClusterAllProxiesPrivate(_ context.Context, _ string) *bool { + return nil +} + func (m *testProxyManager) ClusterSupportsSessionCode(_ context.Context, _ string) bool { return m.supportsSessionCode }