mirror of
https://github.com/netbirdio/netbird.git
synced 2026-09-29 18:19:07 +02:00
[management,proxy] Use single-use codes for OIDC session handoff (#7635)
* Generalize PKCE verifier store into SingleUseStore * Generalize PKCE verifier store into SingleUseStore * Extend single-use store to generate one-time retrieval codes * Hand off proxy OIDC session via one-time code instead of URL token * Use the single-use store in integration tests * Read active proxy versions by cluster * Detect proxy clusters that support session codes * Bind OIDC session handoff mode to signed state * Deprecate legacy OIDC session token handoff * Remove unrelated session code test stub * fix tests * fix merge * Fix session code compatibility detection * Isolate proxy session codes in shared cache * bump min session version
This commit is contained in:
@@ -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
|
||||
ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool
|
||||
CleanupStale(ctx context.Context, inactivityDuration time.Duration) error
|
||||
GetAccountProxy(ctx context.Context, accountID string) (*Proxy, error)
|
||||
CountAccountProxies(ctx context.Context, accountID string) (int64, error)
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"go.opentelemetry.io/otel/metric"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
nbversion "github.com/netbirdio/netbird/version"
|
||||
)
|
||||
|
||||
// store defines the interface for proxy persistence operations
|
||||
@@ -22,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
|
||||
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)
|
||||
CountProxiesByAccountID(ctx context.Context, accountID string) (int64, error)
|
||||
@@ -29,6 +31,8 @@ type store interface {
|
||||
DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error
|
||||
}
|
||||
|
||||
const minSessionCodeVersion = "0.81.0"
|
||||
|
||||
// Manager handles all proxy operations
|
||||
type Manager struct {
|
||||
store store
|
||||
@@ -145,6 +149,22 @@ func (m Manager) ClusterSupportsPrivate(ctx context.Context, clusterAddr string)
|
||||
return m.store.GetClusterSupportsPrivate(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)
|
||||
if err != nil || len(versions) == 0 {
|
||||
return false
|
||||
}
|
||||
|
||||
for _, version := range versions {
|
||||
if supported, err := nbversion.MeetsMinVersion(minSessionCodeVersion, version); err != nil || !supported {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
// CleanupStale removes proxies that haven't sent heartbeat in the specified duration
|
||||
func (m *Manager) CleanupStale(ctx context.Context, inactivityDuration time.Duration) error {
|
||||
if err := m.store.CleanupStaleProxies(ctx, inactivityDuration); err != nil {
|
||||
|
||||
@@ -22,6 +22,7 @@ type mockStore struct {
|
||||
updateProxyHeartbeatFunc func(ctx context.Context, p *proxy.Proxy) error
|
||||
getActiveProxyClusterAddressesFunc func(ctx context.Context) ([]string, error)
|
||||
getActiveProxyClusterAddressesForAccFunc func(ctx context.Context, accountID string) ([]string, error)
|
||||
getActiveProxyVersionsFunc func(ctx context.Context, clusterAddress string) ([]string, error)
|
||||
cleanupStaleProxiesFunc func(ctx context.Context, d time.Duration) error
|
||||
getProxyByAccountIDFunc func(ctx context.Context, accountID string) (*proxy.Proxy, error)
|
||||
countProxiesByAccountIDFunc func(ctx context.Context, accountID string) (int64, error)
|
||||
@@ -104,6 +105,12 @@ func (m *mockStore) GetClusterSupportsCrowdSec(_ context.Context, _ string) *boo
|
||||
func (m *mockStore) GetClusterSupportsPrivate(_ 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)
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func newTestManager(s store) *Manager {
|
||||
meter := noop.NewMeterProvider().Meter("test")
|
||||
@@ -114,6 +121,34 @@ func newTestManager(s store) *Manager {
|
||||
return m
|
||||
}
|
||||
|
||||
func TestClusterSupportsSessionCode(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
versions []string
|
||||
storeErr error
|
||||
want bool
|
||||
}{
|
||||
{name: "all supported", versions: []string{"0.81.0", "0.81.2"}, want: true},
|
||||
{name: "one old proxy", versions: []string{"0.81.0", "0.80.0"}},
|
||||
{name: "missing version", versions: []string{"0.81.0", ""}},
|
||||
{name: "no active proxies"},
|
||||
{name: "store error", storeErr: errors.New("db error")},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
s := &mockStore{
|
||||
getActiveProxyVersionsFunc: func(_ context.Context, _ string) ([]string, error) {
|
||||
return tt.versions, tt.storeErr
|
||||
},
|
||||
}
|
||||
|
||||
got := newTestManager(s).ClusterSupportsSessionCode(context.Background(), "cluster.example.com")
|
||||
assert.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestConnect_WithAccountID(t *testing.T) {
|
||||
accountID := "acc-123"
|
||||
|
||||
|
||||
@@ -112,6 +112,20 @@ func (mr *MockManagerMockRecorder) ClusterSupportsPrivate(ctx, clusterAddr any)
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsPrivate", reflect.TypeOf((*MockManager)(nil).ClusterSupportsPrivate), ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// ClusterSupportsSessionCode mocks base method.
|
||||
func (m *MockManager) ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "ClusterSupportsSessionCode", ctx, clusterAddr)
|
||||
ret0, _ := ret[0].(bool)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// ClusterSupportsSessionCode indicates an expected call of ClusterSupportsSessionCode.
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsSessionCode(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsSessionCode", reflect.TypeOf((*MockManager)(nil).ClusterSupportsSessionCode), ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// Connect mocks base method.
|
||||
func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *Capabilities) (*Proxy, error) {
|
||||
m.ctrl.T.Helper()
|
||||
|
||||
Reference in New Issue
Block a user