[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:
Bethuel Mmbaga
2026-09-29 18:29:55 +03:00
committed by GitHub
parent 7ff709f565
commit 30dd076b36
22 changed files with 676 additions and 281 deletions
+12 -3
View File
@@ -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 {
+1
View File
@@ -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)
+15
View File
@@ -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()