mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-01 02:59:08 +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:
@@ -59,7 +59,7 @@ func (h *AuthCallbackHandler) handleCallback(w http.ResponseWriter, r *http.Requ
|
||||
|
||||
state := r.URL.Query().Get("state")
|
||||
|
||||
codeVerifier, originalURL, err := h.proxyService.ValidateState(state)
|
||||
codeVerifier, originalURL, useSessionCode, err := h.proxyService.ValidateState(state)
|
||||
if err != nil {
|
||||
log.WithError(err).Error("OAuth callback state validation failed")
|
||||
http.Error(w, "Invalid state parameter", http.StatusBadRequest)
|
||||
@@ -119,10 +119,19 @@ func (h *AuthCallbackHandler) handleCallback(w http.ResponseWriter, r *http.Requ
|
||||
redirectURL.Scheme = "https"
|
||||
|
||||
query := redirectURL.Query()
|
||||
query.Set("session_token", sessionToken)
|
||||
if useSessionCode {
|
||||
code, ok := h.proxyService.GenerateSessionCode(sessionToken)
|
||||
if !ok {
|
||||
http.Error(w, "Failed to create session", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
query.Set("session_code", code)
|
||||
} else {
|
||||
query.Set("session_token", sessionToken)
|
||||
}
|
||||
redirectURL.RawQuery = query.Encode()
|
||||
|
||||
log.WithField("redirect", redirectURL.Host).Debug("OAuth callback: redirecting user with session token")
|
||||
log.WithField("redirect", redirectURL.Host).Debug("OAuth callback: redirecting user to proxy")
|
||||
http.Redirect(w, r, redirectURL.String(), http.StatusFound)
|
||||
}
|
||||
|
||||
|
||||
@@ -181,6 +181,10 @@ func (m *testAccessLogManager) GetAllAccessLogs(_ context.Context, _, _ string,
|
||||
}
|
||||
|
||||
func setupAuthCallbackTest(t *testing.T) *testSetup {
|
||||
return setupAuthCallbackTestWithProxyManager(t, testSessionCodeManager{})
|
||||
}
|
||||
|
||||
func setupAuthCallbackTestWithProxyManager(t *testing.T, proxyManager nbproxy.Manager) *testSetup {
|
||||
t.Helper()
|
||||
|
||||
ctx := context.Background()
|
||||
@@ -217,7 +221,7 @@ func setupAuthCallbackTest(t *testing.T) *testSetup {
|
||||
nil,
|
||||
usersManager,
|
||||
nil,
|
||||
nil,
|
||||
proxyManager,
|
||||
nil,
|
||||
)
|
||||
|
||||
@@ -242,6 +246,15 @@ func setupAuthCallbackTest(t *testing.T) *testSetup {
|
||||
}
|
||||
}
|
||||
|
||||
type testSessionCodeManager struct {
|
||||
nbproxy.Manager
|
||||
supported bool
|
||||
}
|
||||
|
||||
func (m testSessionCodeManager) ClusterSupportsSessionCode(_ context.Context, _ string) bool {
|
||||
return m.supported
|
||||
}
|
||||
|
||||
func createTestReverseProxies(t *testing.T, ctx context.Context, testStore store.Store) {
|
||||
t.Helper()
|
||||
|
||||
@@ -252,10 +265,11 @@ func createTestReverseProxies(t *testing.T, ctx context.Context, testStore store
|
||||
privKey := base64.StdEncoding.EncodeToString(priv)
|
||||
|
||||
testProxy := &service.Service{
|
||||
ID: "testProxyId",
|
||||
AccountID: "testAccountId",
|
||||
Name: "Test Proxy",
|
||||
Domain: "test-proxy.example.com",
|
||||
ID: "testProxyId",
|
||||
AccountID: "testAccountId",
|
||||
Name: "Test Proxy",
|
||||
Domain: "test-proxy.example.com",
|
||||
ProxyCluster: "cluster.example.com",
|
||||
Targets: []*service.Target{{
|
||||
Path: strPtr("/"),
|
||||
Host: "localhost",
|
||||
@@ -512,29 +526,56 @@ func createTestState(t *testing.T, ps *nbgrpc.ProxyServiceServer, redirectURL st
|
||||
}
|
||||
|
||||
func TestAuthCallback_UserAllowedToLogin(t *testing.T) {
|
||||
setup := setupAuthCallbackTest(t)
|
||||
defer setup.cleanup()
|
||||
tests := []struct {
|
||||
name string
|
||||
manager nbproxy.Manager
|
||||
wantParam string
|
||||
absentParam string
|
||||
}{
|
||||
{name: "legacy proxy", manager: testSessionCodeManager{}, wantParam: "session_token", absentParam: "session_code"},
|
||||
{name: "compatible proxy", manager: testSessionCodeManager{supported: true}, wantParam: "session_code", absentParam: "session_token"},
|
||||
}
|
||||
|
||||
setup.oidcServer.tokenSubject = "allowedUserId"
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
setup := setupAuthCallbackTestWithProxyManager(t, tt.manager)
|
||||
defer setup.cleanup()
|
||||
|
||||
state := createTestState(t, setup.proxyService, "https://test-proxy.example.com/dashboard")
|
||||
setup.oidcServer.tokenSubject = "allowedUserId"
|
||||
state := createTestState(t, setup.proxyService, "https://test-proxy.example.com/dashboard")
|
||||
req := httptest.NewRequest(http.MethodGet, "/reverse-proxy/callback?code=test-auth-code&state="+url.QueryEscape(state), nil)
|
||||
rec := httptest.NewRecorder()
|
||||
setup.router.ServeHTTP(rec, req)
|
||||
require.Equal(t, http.StatusFound, rec.Code)
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/reverse-proxy/callback?code=test-auth-code&state="+url.QueryEscape(state), nil)
|
||||
rec := httptest.NewRecorder()
|
||||
location, err := url.Parse(rec.Header().Get("Location"))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "test-proxy.example.com", location.Host)
|
||||
require.NotEmpty(t, location.Query().Get(tt.wantParam))
|
||||
require.Empty(t, location.Query().Get(tt.absentParam))
|
||||
require.Empty(t, location.Query().Get("error"))
|
||||
|
||||
setup.router.ServeHTTP(rec, req)
|
||||
if tt.wantParam == "session_code" {
|
||||
code := location.Query().Get("session_code")
|
||||
response, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
|
||||
Domain: location.Hostname(),
|
||||
SessionCode: code,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.True(t, response.GetValid())
|
||||
require.NotEmpty(t, response.GetSessionToken())
|
||||
require.NotEqual(t, code, response.GetSessionToken())
|
||||
|
||||
require.Equal(t, http.StatusFound, rec.Code)
|
||||
|
||||
location := rec.Header().Get("Location")
|
||||
require.NotEmpty(t, location)
|
||||
|
||||
parsedLocation, err := url.Parse(location)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Equal(t, "test-proxy.example.com", parsedLocation.Host)
|
||||
require.NotEmpty(t, parsedLocation.Query().Get("session_token"), "Should include session token")
|
||||
require.Empty(t, parsedLocation.Query().Get("error"), "Should not have error parameter")
|
||||
replayed, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
|
||||
Domain: location.Hostname(),
|
||||
SessionCode: code,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.False(t, replayed.GetValid())
|
||||
require.Empty(t, replayed.GetSessionToken())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestAuthCallback_UserDeniedByAccountStatus asserts that a user whose account
|
||||
|
||||
@@ -367,6 +367,21 @@ func (s *SqlStore) GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr s
|
||||
return s.getClusterUnanimousCapability(ctx, clusterAddr, "supports_crowdsec")
|
||||
}
|
||||
|
||||
// GetActiveProxyVersions returns every active proxy version in a cluster.
|
||||
func (s *SqlStore) GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error) {
|
||||
var versions []string
|
||||
err := s.db.WithContext(ctx).
|
||||
Model(&proxy.Proxy{}).
|
||||
Where("LOWER(cluster_address) = LOWER(?) AND status = ? AND last_seen > ?",
|
||||
clusterAddr, proxy.StatusConnected, time.Now().Add(-proxyActiveThreshold)).
|
||||
Pluck("version", &versions).Error
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("failed to get active proxy versions for %s: %v", clusterAddr, err)
|
||||
return nil, status.Errorf(status.Internal, "get active proxy versions")
|
||||
}
|
||||
return versions, nil
|
||||
}
|
||||
|
||||
// getClusterUnanimousCapability returns an aggregated boolean capability
|
||||
// requiring all active proxies in the cluster to report true.
|
||||
func (s *SqlStore) getClusterUnanimousCapability(ctx context.Context, clusterAddr, column string) *bool {
|
||||
|
||||
@@ -337,6 +337,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
|
||||
GetAllProxies(ctx context.Context) ([]*proxy.Proxy, error)
|
||||
DisconnectAllProxies(ctx context.Context) (int64, error)
|
||||
|
||||
@@ -1539,6 +1539,21 @@ func (mr *MockStoreMockRecorder) GetActiveProxyClusterAddressesForAccount(ctx, a
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveProxyClusterAddressesForAccount", reflect.TypeOf((*MockStore)(nil).GetActiveProxyClusterAddressesForAccount), ctx, accountID)
|
||||
}
|
||||
|
||||
// GetActiveProxyVersions mocks base method.
|
||||
func (m *MockStore) GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetActiveProxyVersions", ctx, clusterAddr)
|
||||
ret0, _ := ret[0].([]string)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetActiveProxyVersions indicates an expected call of GetActiveProxyVersions.
|
||||
func (mr *MockStoreMockRecorder) GetActiveProxyVersions(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveProxyVersions", reflect.TypeOf((*MockStore)(nil).GetActiveProxyVersions), ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// GetAgentNetworkAccessLogSessions mocks base method.
|
||||
func (m *MockStore) GetAgentNetworkAccessLogSessions(ctx context.Context, lockStrength LockingStrength, accountID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLogSession, int64, error) {
|
||||
m.ctrl.T.Helper()
|
||||
|
||||
Reference in New Issue
Block a user